Compare commits

..
Author SHA1 Message Date
shadowcz007 3191cf2af0 test 2023-12-29 13:27:48 +08:00
shadowcz007 7d3e1ae945 test 2023-12-29 12:59:27 +08:00
shadowcz007 bfced5b4da test 2023-12-29 12:49:05 +08:00
174 changed files with 4257 additions and 45729 deletions
-21
View File
@@ -1,21 +0,0 @@
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 -4
View File
@@ -2,7 +2,4 @@ __pycache__/
https/
nodes/config.json
workflow/my_workflow.json
workflow/my_workflow_app.json
workflow/prompt_result.json
app/*
workflow/prompt_result.json
workflow/my_workflow_app.json
-1
View File
@@ -1 +0,0 @@
mixlabnodes.com
-21
View File
@@ -1,21 +0,0 @@
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.
+54 -208
View File
@@ -1,263 +1,114 @@
![](https://img.shields.io/github/release/shadowcz007/comfyui-mixlab-nodes)
> 适配了最新版 comfyui 的 py3.11 ,torch 2.1.2+cu121
> [Mixlab nodes discord](https://discord.gg/cXs9vZSqeK)
##### `最新`:
- 增加 SiliconflowLLM,可以使用由Siliconflow提供的免费LLM
- 增加 Edit Mask,方便在生成的时候手动绘制 mask [workflow](./workflow/edit-mask-workflow.json)
<!-- - 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) -->
## 🚀🚗🚚🏃 Workflow-to-APP
- 新增 AppInfo 节点,可以通过简单的配置,把 workflow 转变为一个 Web APP。
- 支持多个 web app 切换
- 发布为 app 的 workflow,可以在右键里再次编辑了
- web app 可以设置分类,在 comfyui 右键菜单可以编辑更新 web app
- 支持动态提示
- 支持把输出显示到comfyui背景(TouchDesigner 风格)
![](./assets/微信图片_20240421205440.png)
- Support multiple web app switching.
##
v0.6.0 🚀🚗🚚🏃‍ Workflow-to-APP
- 新增AppInfo节点,可以通过简单的配置,把workflow转变为一个Web APP。
- 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)
![](./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)
APP-JSON:
- [text-to-image](./example/Text-to-Image_3.json)
- [image-to-image](./example/Image-to-Image_2.json)
- [text-to-image](./app/text-to-image_1_Wed%20Dec%2027%202023.json)
- [image-to-image](./app/image-to-image_1_Wed%20Dec%2027%202023.json)
- text-to-text
> 暂时支持 9 种节点作为界面上的输入节点:Load Image、VHS*LoadVideo、CLIPTextEncode、PromptSlide、TextInput*、Color、FloatSlider、IntNumber、CheckpointLoaderSimple、LoraLoader
> 暂时支持6种节点作为界面上的输入节点:Load Image、CLIPTextEncode、TextInput_、FloatSlider、CheckpointLoaderSimple、LoraLoader
> 输出节点:PreviewImage 、SaveImage、ShowTextForGPT、VHS_VideoCombine、PromptImage
> 输出节点:PreviewImage 、SaveImage、ShowTextForGPT
> seed 统一输入控件,支持:SamplerCustom、KSampler
> 配套[ps 插件](https://github.com/shadowcz007/comfyui-ps-plugin)
### 3D
![](./assets/3dimage.png)
[workflow](./workflow/3D-workflow.json)
> 如果遇到上传图片不成功,请检查下:局域网或者是云服务,请使用 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
> 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! 💻🌐
### 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](./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)
### GPT
> Support for calling multiple GPTs.ChatGPT、ChatGLM3 , 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
> Support for calling multiple GPTs.Local LLM(llama.cpp)、 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)
[workflow-5](./workflow/5-gpt-workflow.json)
最新: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)
<!-- ![](./workflow/promptslide-appinfo-workflow.svg) -->
> randomPrompt
![randomPrompt](./assets/randomPrompt.png)
> ClipInterrogator
[add clip-interrogator](https://github.com/pharmapsychotic/clip-interrogator)
> PromptImage & PromptSimplification,Assist in simplifying prompt words, comparing images and prompt word nodes.
> 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.
#### 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.
### 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.
![watch](./assets/4-loadfromlocal-watcher-workflow.svg)
[workflow-4](./workflow/4-loadfromlocal-watcher-workflow.json)
#### LoadImagesFromURL
### 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```
### 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.
![layers](./assets/layers-workflow.svg)
![poster](./assets/poster-workflow.svg)
### Style
> 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)
### Utils
## 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 compare the effects of different models.](./workflow/ckpts-image-workflow.json)
### Other Nodes
## Other Nodes
![main](./assets/all-workflow.svg)
![main2](./assets/detect-face-all.png)
[workflow-1](./workflow/1-workflow.json)
> randomPrompt
![randomPrompt](./assets/randomPrompt.png)
> TransparentImage
![TransparentImage](./assets/TransparentImage.png)
> Consistency Decoder
[openai Consistency Decoder]( https://github.com/openai/consistencydecoder)
![Consistency](./assets/consistency.png)
After downloading the OpenAI VAE model, place it in the "model/vae" directory for use.
https://openaipublic.azureedge.net/diff-vae/c9cebd3132dd9c42936d803e33424145a748843c8f716c0814838bdc8a2fe7cb/decoder.pt
> FeatheredMask、SmoothMask
Add edges to an image.
![FeatheredMask](./assets/FlVou_Y6kaGWYoEj1Tn0aTd4AjMI.jpg)
> LaMaInpainting
from [simple-lama-inpainting](https://github.com/enesmsahin/simple-lama-inpainting)
> 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
### Improvement
### Improvement
- Add "help" option to the context menu for each node.
- Add "Nodes Map" option to the global context menu.
@@ -268,21 +119,14 @@ An improvement has been made to directly redirect to GitHub to search for missin
![node-not-found](./assets/node-not-found.png)
### Models
[Download CLIPSeg](https://huggingface.co/CIDAS/clipseg-rd64-refined/tree/main), move to : model/clipseg
- [Download TripoSR](https://huggingface.co/stabilityai/TripoSR/blob/main/model.ckpt) and place it in `models/triposr`
<!-- ### Workflow
[Workflow](./workflow.md) -->
- [Download facebook/dino-vitb16](https://huggingface.co/facebook/dino-vitb16/tree/main) and place it in `models/triposr/facebook/dino-vitb16`
[Download rembg Models](https://github.com/danielgatis/rembg/tree/main#Models),move to:`models/rembg`
[Download lama](https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt), move to : `models/lama`
[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
@@ -298,36 +142,37 @@ 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
#### Thanks:
[ComfyUI-CLIPSeg](https://github.com/biegert/ComfyUI-CLIPSeg/tree/main)
#### discussions:
[discussions](https://github.com/shadowcz007/comfyui-mixlab-nodes/discussions)
### TODO:
- 音频播放节点:带可视化、支持多音轨、可配置音轨音量
- vector https://github.com/GeorgLegato/stable-diffusion-webui-vectorstudio
<picture>
<source
media="(prefers-color-scheme: dark)"
@@ -346,3 +191,4 @@ File / LoadImagesFromPath SaveImageToLocal LoadImagesFromURL
src="https://api.star-history.com/svg?repos=shadowcz007/comfyui-mixlab-nodes&type=Date"
/>
</picture>
+102 -927
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
Binary file not shown.

Before

Width:  |  Height:  |  Size: 101 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 135 KiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 210 KiB

File diff suppressed because one or more lines are too long

Before

Width:  |  Height:  |  Size: 2.4 MiB

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

Before

Width:  |  Height:  |  Size: 1.1 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 11 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 784 KiB

Binary file not shown.
Binary file not shown.
Binary file not shown.

Before

Width:  |  Height:  |  Size: 75 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 63 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 477 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 965 KiB

-30
View File
@@ -1,30 +0,0 @@
Jony Ive
Dieter Rams
Philippe Starck
Karim Rashid
Yves Béhar
Marc Newson
Naoto Fukasawa
Jonathan Adler
Patricia Urquiola
Ross Lovegrove
Tom Dixon
Jasper Morrison
Charles Eames
Ray Eames
Achille Castiglioni
Ron Arad
Konstantin Grcic
Marcel Wanders
Maarten Baas
Stefan Sagmeister
Ingo Maurer
Hella Jongerius
Sam Hecht
Kim Colin
Jaime Hayon
Michael Anastassiades
Nendo
Oki Sato
Matali Crasset
Tokujin Yoshioka
-10
View File
@@ -1,10 +0,0 @@
Chibi Anime Style
Gakuen Anime Style
Gekiga Anime Style
Jidaimono Anime Style
Kawaii Anime Style
Mecha Anime Style
Realistic Anime Style
Semi-Realistic Anime Style
Shoji Anime Style
Kemonomimi Anime Style
-2052
View File
File diff suppressed because it is too large Load Diff
-23
View File
@@ -1,23 +0,0 @@
GoPro
Drone
polaroid
black and white film
Kodachrome
shot on 8mm
shot on 16mm
shot on 35mm
Microscopic
Fisheye Lens
Wide Angle
Ultra-Wide Angle
Panorama
Short Exposure
Long Exposure
Double Exposure
f2.8
Depth of Field
Soft Focus
Deep Focus
Shallow Focus
Vanishing Point
Vantage Point
-30
View File
@@ -1,30 +0,0 @@
Elegant evening gown
Casual jeans and t-shirt
Formal black suit
Stylish leather jacket
Flowy bohemian dress
Sporty tracksuit
Chic little black dress
Trendy ripped jeans
Classic white button-down shirt
Cozy oversized sweater
Sophisticated tailored blazer
Quirky patterned leggings
Striped sailor top
Polished knee-length skirt
Vintage-inspired floral dress
Edgy motorcycle jacket
Preppy polo shirt
Boho maxi skirt
Professional pinstripe suit
Relaxed denim shorts
Glamorous sequined dress
Athletic running shoes
Formal bow tie
Casual baseball cap
Stylish fedora hat
Warm woolen scarf
Comfortable cotton socks
Trendy ankle boots
Cute summer sandals
Cozy pajama set
-30
View File
@@ -1,30 +0,0 @@
Happy
Sad
Angry
Surprised
Excited
Worried
Confused
Disgusted
Amused
Bored
Curious
Embarrassed
Frustrated
Nervous
Pleased
Relieved
Shy
Tired
Serious
Silly
Proud
Grumpy
Smug
Sarcastic
Flirty
Skeptical
Shocked
Blissful
Envious
Mischievous
+493 -9213
View File
File diff suppressed because it is too large Load Diff
-10
View File
@@ -1,10 +0,0 @@
[
{
"keyword":"Dog",
"imgurl":"http://127.0.0.1:8188/view?filename=1709966910233.png&type=input&subfolder=&rand=0.2734446552394221"
},
{
"keyword":"x",
"imgurl":"http://127.0.0.1:8188/view?filename=image%20(33).png&type=input&subfolder=pasted&rand=0.6984318219852814"
}
]
-16
View File
@@ -1,16 +0,0 @@
Mood Lighting
Moody Lighting
Studio Lighting
Cove Lighting
Soft Lighting
Hard Lighting
Volumetric Lighting
Low-Key Lighting
High-Key Lighting
Epic Light
Rembrandt Lighting
Contre-Jour
Veiling Flare
Crepuscular Rays
Rays of Shimmering Light
Godrays
-1
View File
@@ -1 +0,0 @@
{}
-132
View File
@@ -1,132 +0,0 @@
Aaron Siskind
Alessio Albi
Alfred Eisenstaedt
Alfred Stieglitz
Alyssa Monks
André Kertész
Andreas Gursky
Andrew Wyeth
Anne Geddes
Annie Leibovitz
Ansel Adams
Arnold Newman
August Sander
Balthus
Berenice Abbott
Bill Brandt
Bill Henson
Brassaï (Gyula Halász)
Brooke Shaden
Bruce Davidson
Bruce Weber
Bunny Yeager
Carleton Watkins
Carrie Mae Weems
Chuck Close
Cindy Sherman
Clarence H. White
Claude Cahun
Danny Lyon
David LaChapelle
Dawoud Bey
Diane Arbus
Don McCullin
Dora Maar
Dorothea Lange
Duane Michals
Eadweard Muybridge
Edward Burtynsky
Edward Curtis
Edward Ruscha
Edward Steichen
Edward Weston
Elliott Erwitt
Ernst Haas
Eugene Atget
Fan Ho
Francesca Woodman
Frans Lanting
Garry Winogrand
Georges Melies
Gerda Taro
Gertrude Käsebier
Gordon Parks
Graciela Iturbide
Gregory Crewdson
Harold Edgerton
Helen Levitt
Helmut Newton
Hendrik Kerstens
Henri Cartier-Bresson
Hugh Kretschmer
Irving Penn
Jacques Henri Lartigue
James Nachtwey
James Van Der Zee
Jay Maisel
Jerry Uelsmann
Joel Peter Witkin
Joel Sartore
John Frederick William Herschel
Josef Sudek
Julia Margaret Cameron
Karl Blossfeldt
Larry Burrows
László Moholy-Nagy (photography)
Lee Jeffries
Lewis Hine
Lorna Simpson
Lynsey Addario
Margaret Bourke-White
Mario Testino
Martin Parr
Martin Schoeller
Mary Ellen Mark
Mathew B. Brady
Méret Oppenheim
Meryl McMaster
Mick Rock
Miles Aldridge
Minor Martin White
Nan Goldin
Nathan Wirth
Olive Cotton
Olivier Rousteing
Patrick Demarchelier
Paul Nicklen
Paul Outerbridge
Paul Strand
Pete Souza
Peter Dombrovskis
Peter Henry Emerson
Peter Lik
Peter Lindbergh
Philip-Lorca diCorcia
Philippe Halsman
Ralph Gibson
Richard Avedon
Robert Adams
Robert Bechtle
Robert Capa
Robert Frank
Robert Mapplethorpe
Roger Fenton
Ruth Bernhard
Sally Mann
Sebastião Salgado
Shirin Neshat
Stefan Gesell
Steven Meisel
Susan Meiselas
Vivian Maier
Vivian Maier
Viviane Sassen
Walker Evans
Wes Anderson
William Eggleston
William Eugene Smith
William Henry Fox Talbot
Yinka Shonibare
Yousuf Karsh
Man Ray
Robert Mapplethorpe
-101
View File
@@ -1,101 +0,0 @@
Doctor
Teacher
Engineer
Lawyer
Accountant
Nurse
Architect
Chef
Pilot
Scientist
Artist
Writer
Musician
Actor
Photographer
Police officer
Firefighter
Dentist
Pharmacist
Veterinarian
Electrician
Plumber
Carpenter
Mechanic
Farmer
Astronaut
Athlete
Journalist
Politician
Economist
Psychologist
Social worker
Librarian
Translator
Salesperson
Entrepreneur
Financial advisor
Graphic designer
Web developer
Marketing manager
Human resources manager
Project manager
Event planner
Fashion designer
Interior decorator
Real estate agent
Archaeologist
Biologist
Chemist
Geologist
Physicist
Mathematician
Historian
Geographer
Economist
Sociologist
Anthropologist
Archaeologist
Linguist
Philosopher
Economist
Sociologist
Anthropologist
Archaeologist
Linguist
Philosopher
Geographer
Historian
Economist
Sociologist
Anthropologist
Archaeologist
Linguist
Philosopher
Geographer
Historian
Economist
Sociologist
Anthropologist
Archaeologist
Linguist
Philosopher
Geographer
Historian
Economist
Sociologist
Anthropologist
Archaeologist
Linguist
Philosopher
Geographer
Historian
Economist
Sociologist
Anthropologist
Archaeologist
Linguist
Philosopher
Geographer
Historian
#MixCopilot
-58
View File
@@ -1,58 +0,0 @@
Residential space
Apartment building
Villa
Bungalow
Condominium
Commercial space
Shopping mall
Supermarket
Restaurant
Store
Market
Office space
Office building
Office
Meeting room
Co-working space
Educational space
School
University
Training institution
Library
Laboratory
Medical space
Hospital
Clinic
Pharmacy
Nursing home
Rehabilitation center
Cultural space
Museum
Library
Theater
Concert hall
Gallery
Sports space
Sports stadium
Gym
Swimming pool
Basketball court
Football field
Transportation space
Airport
Train station
Subway station
Bus stop
Parking lot
Public space
Park
Square
Street
Pedestrian street
Community center
Industrial space
Factory
Warehouse
Production workshop
Mine
Power plant
-135
View File
@@ -1,135 +0,0 @@
Vintage
Grain
Sepia
High Key
Low Key
High Dynamic Range
Cross Process
Radial Blur
Infrared
Lomo
Photocopy
Pencil Sketch
Pop Art
Orton
Mosaic
Selective Black and White
Torn Paper
Tilt-Shift
Double Exposure
Polaroid
Liquid Ink
Color Splash
Sketch
Water Drops
Polarizer
Chinese Painting
Water Droplets
Polarization
Color Inversion
Fish-eye
Soft Focus
Solarization
Posterize
Comic Book
Duotone
Gradient Map
Edge Detection
Oil Painting
Reflection
Mirror
ASCII Art
Glitch
Time-Lapse
Day to Night
Surreal
Black and White
Sepia Tone
Vintage Film
Grainy Texture
High Key Lighting
Low Key Lighting
Cross Processed Film
Infrared Photography
Photocopy
Pencil Drawing
Pop Art Filter
Mosaic Filter
Selective Desaturation
Torn Paper
Tilt-Shift Photography
Double Exposure
Polaroid Style Frame
Water Drops Texture
Polarizer
Chinese Painting
Water Droplets Texture
Polarization
Color Inversion
Fish-eye Lens
Soft Focus
Solarize Filter
Edge Detection
Oil Painting
Reflection
Mirror Image
Time-Lapse Photography
Day to Night Transition
Surreal Art Style
Abstract Expressionism
Acrylic Painting
Anime
Art Deco
Biomorphic Abstraction
Black and White Photograph
Cartoon
Charcoal Sketch
Chibi Anime
Chinese Painting
Classicist Painting
Collage
Concept Art
Cyberpunk
Dada Art
Digital Art
Fantasy Art
Fashion Art
Fashion Sketch
Fish-Eye lens Photograph
Goth Art
Graffiti
Harlem Renaissance
High Key Photograph
Hyperrealist Pencil Sketch
Impressionist Painting
Josei Anime
Long Exposure Photograph
Low Key Photograph
Macro Photograph
Manga
Metal Sculpture
Mid Century Modern Illustration
Mixed Media
Modern Art
Moe Anime
Nihonga
Origami
Paper Mache
Pen and Ink
Pencil Sketch
Photograph
Photorealism
Pinup Art
Romanticist Painting
Sci-Fi Art
Semi Realistic Fantasy Art
Semi Realistic Cyberpunk Art
Shallow Depth of Field Photograph
Steam Punk Art
Stone Sculpture
Superhero Comic
Surrealist Art
Tempura Painting
Underground Comic
Watercolor Painting
Zulu Urban Art
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,12 +10,6 @@ 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
)
%python_exec% -s -m pip install --upgrade --force llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
%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 (
+39 -59
View File
@@ -1,7 +1,6 @@
import os
import folder_paths
import torchaudio
class SpeechRecognition:
@classmethod
@@ -25,7 +24,7 @@ class SpeechRecognition:
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Audio"
CATEGORY = "♾️Mixlab/audio"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
@@ -49,72 +48,53 @@ class SpeechSynthesis:
OUTPUT_NODE = True
OUTPUT_IS_LIST = (True,)
CATEGORY = "♾️Mixlab/Audio"
CATEGORY = "♾️Mixlab/audio"
def run(self, text):
# print(session_history)
return {"ui": {"text": text}, "result": (text,)}
class AudioPlayNode:
def __init__(self):
self.output_dir = folder_paths.get_temp_directory()
self.type = "temp"
self.prefix_append = ""
self.compress_level = 4
#
class GamePal:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"audio": ("AUDIO",),
},
}
RETURN_TYPES = ()
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Audio"
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
OUTPUT_IS_LIST = ()
RETURN_TYPES = ("INT",)
FUNCTION = "run"
OUTPUT_NODE = True
def run(self,audio):
OUTPUT_IS_LIST = (False,)
# 判断是否是 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
CATEGORY = "♾️Mixlab/audio"
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']
}]
def run(self, input_text,input_num,python_code):
exec(python_code)
res=None
try:
# 可能会引发异常的代码
res=result
except:
# 处理异常的代码
print('')
# print(audio)
return {"ui": {"audio":results}}
print(res)
# print(session_history)
return {"ui": {"text": [input_text],"num":[input_num]}, "result": (res,)}
+28 -400
View File
@@ -1,37 +1,7 @@
import openai
import time
import urllib.error
import re,json,os,string,random
import folder_paths
import hashlib
import codecs,sys
import importlib.util
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())
unique_hash = hash_object.hexdigest()
return unique_hash
def generate_random_string(length):
letters = string.ascii_letters + string.digits
return ''.join(random.choice(letters) for _ in range(length))
class AnyType(str):
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
def __ne__(self, __value: object) -> bool:
return False
any_type = AnyType("*")
import re,json
# 判断是否是azure服务
def is_azure_url(url):
@@ -53,136 +23,23 @@ 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')==False:
import subprocess
# 安装
print('#pip install zhipuai')
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'zhipuai'], capture_output=True, text=True)
#检查命令执行结果
if result.returncode == 0:
print("#install success")
from zhipuai import ZhipuAI
else:
print("#install error")
else:
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()
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
def chat(client, model_name,messages ):
print('#chat',model_name,messages)
try_count = 0
while True:
try_count += 1
try:
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,
)
response = client.chat.completions.create(
model=model_name,
messages=messages
)
break
except openai.AuthenticationError as ex:
raise ex
@@ -191,8 +48,7 @@ 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)
@@ -215,45 +71,18 @@ class ChatGPTNode:
@classmethod
def INPUT_TYPES(cls):
model_list=llama_modes_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}),
"prompt": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"api_key":("KEY", {"default": "", "multiline": True}),
"api_url":("URL", {"default": "", "multiline": True}),
"prompt": ("STRING", {"multiline": True}),
"system_content": ("STRING",
{
"default": "You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible.",
"multiline": True,"dynamicPrompts": False
"multiline": True
}),
"model": ( model_list,
{"default": model_list[0]}),
"model": (["gpt-3.5-turbo","gpt-35-turbo","gpt-3.5-turbo-16k", "gpt-3.5-turbo-16k-0613", "gpt-4-0613","gpt-4-1106-preview"],
{"default": "gpt-3.5-turbo"}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
"context_size":("INT", {"default": 1, "min": 0, "max":30, "step": 1}),
},
@@ -276,8 +105,8 @@ class ChatGPTNode:
api_url,
prompt,
system_content,
model,
seed,context_size,unique_id = None, extra_pnginfo=None):
model,
seed,context_size,unique_id = None, extra_pnginfo=None):
# print(api_key!='',api_url,prompt,system_content,model,seed)
# 可以选择保留会话历史以维持上下文记忆
# 或者在此处清除会话历史 self.session_history.clear()
@@ -295,16 +124,8 @@ class ChatGPTNode:
if is_azure_url(api_url):
client=azure_client(api_key,api_url)
else:
# 根据用户选择的模型,设置相应的接口和模型名称
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',api_key,api_url)
client=openai_client(api_key,api_url)
print('openai url')
# 把用户的提示添加到会话历史中
# 调用API时传递整个会话历史
@@ -320,7 +141,6 @@ 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}]
@@ -341,102 +161,14 @@ 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":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
"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}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
"extra_pnginfo": "EXTRA_PNGINFO",
},
}
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,unique_id = None, extra_pnginfo=None):
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
def INPUT_TYPES(s):
return {
"required": {
"text": ("STRING", {"forceInput": True,"dynamicPrompts": False}),
},
"optional":{
"output_dir": ("STRING",{"forceInput": True,"default": "","multiline": True,"dynamicPrompts": False}),
}
"text": ("STRING", {"forceInput": True}),
}
}
INPUT_IS_LIST = True
@@ -445,65 +177,11 @@ class ShowTextForGPT:
OUTPUT_NODE = True
OUTPUT_IS_LIST = (True,)
CATEGORY = "♾️Mixlab/Text"
CATEGORY = "♾️Mixlab/GPT"
def run(self, text,output_dir=[""]):
# 类型纠正
texts=[]
for t in text:
if not isinstance(t, str):
t = str(t)
texts.append(t)
text=texts
if len(output_dir)==1 and (output_dir[0]=='' or os.path.dirname(output_dir[0])==''):
t='\n'.join(text)
output_dir=[
os.path.join(folder_paths.get_temp_directory(),
get_unique_hash(t)+'.txt'
)
]
elif len(output_dir)==1:
base=os.path.basename(output_dir[0])
t='\n'.join(text)
if base=='' or os.path.splitext(base)[1]=='':
base=get_unique_hash(t)+'.txt'
output_dir=[
os.path.join(output_dir[0],
base
)
]
# elif len(output_dir)>1:
if len(output_dir)==1 and len(text)>1:
output_dir=[output_dir[0] for _ in range(len(text))]
for i in range(len(text)):
o_fp=output_dir[i]
dirp=os.path.dirname(o_fp)
if dirp=='':
dirp=folder_paths.get_temp_directory()
o_fp=os.path.join(folder_paths.get_temp_directory(),o_fp
)
if not os.path.exists(dirp):
os.mkdir(dirp)
if not os.path.splitext(o_fp)[1].lower()=='.txt':
o_fp=o_fp+'.txt'
t=text[i]
with open(o_fp, 'w') as file:
file.write(t)
# print(text)
def run(self, text):
# print(session_history)
return {"ui": {"text": text}, "result": (text,)}
class CharacterInText:
@@ -511,8 +189,8 @@ class CharacterInText:
def INPUT_TYPES(s):
return {
"required": {
"text": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"character": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"text": ("STRING", {"multiline": True}),
"character": ("STRING", {"multiline": True}),
"start_index": ("INT", {
"default": 1,
"min": 0, #Minimum value
@@ -529,61 +207,11 @@ class CharacterInText:
# OUTPUT_NODE = True
OUTPUT_IS_LIST = (False,)
CATEGORY = "♾️Mixlab/Text"
CATEGORY = "♾️Mixlab/GPT"
def run(self, text,character,start_index):
# print(text,character,start_index)
b=1 if character.lower() in text.lower() else 0
b=1 if character in text else 0
return (b+start_index,)
class TextSplitByDelimiter:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"text": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"delimiter":("STRING", {"multiline": False,"default":",","dynamicPrompts": False}),
"start_index": ("INT", {
"default": 0,
"min": 0, #Minimum value
"max": 1000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"skip_every": ("INT", {
"default": 0,
"min": 0, #Minimum value
"max": 10, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"max_count": ("INT", {
"default": 10,
"min": 1, #Minimum value
"max": 1000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
}
}
INPUT_IS_LIST = False
RETURN_TYPES = ("STRING",)
FUNCTION = "run"
# OUTPUT_NODE = True
OUTPUT_IS_LIST = (True,)
CATEGORY = "♾️Mixlab/Text"
def run(self, text,delimiter,start_index,skip_every,max_count):
if delimiter=="":
arr=[text.strip()]
else:
delimiter=codecs.decode(delimiter, 'unicode_escape')
arr= [line for line in text.split(delimiter) if line.strip()]
arr= arr[start_index:start_index + max_count * (skip_every+1):(skip_every+1)]
return (arr,)
-283
View File
@@ -1,283 +0,0 @@
import os,sys
import folder_paths
from PIL import Image
import importlib.util
import comfy.utils
import numpy as np
import json
import torch
import random
# from clip_interrogator import Config, Interrogator
global _available
_available=False
def is_installed(package):
try:
spec = importlib.util.find_spec(package)
except ModuleNotFoundError:
return False
return spec is not None
try:
if is_installed('clip_interrogator')==False:
import subprocess
# 安装
print('#pip install clip-interrogator==0.6.0')
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'clip-interrogator==0.6.0'], capture_output=True, text=True)
#检查命令执行结果
if result.returncode == 0:
print("#install success")
from clip_interrogator import Config, Interrogator
_available=True
else:
print("#install error")
else:
from clip_interrogator import Config, Interrogator
_available=True
except:
_available=False
try:
from transformers import AutoProcessor, BlipForConditionalGeneration
except:
_available=False
print('pls check transformers.__version__>=4.36.0:: AutoProcessor, BlipForConditionalGeneration')
def load_caption_model(model_path,config,t='blip-base'):
dtype=torch.float16 if config.device == 'cuda' else torch.float32
caption_model = BlipForConditionalGeneration.from_pretrained(model_path, torch_dtype=dtype)
caption_processor = AutoProcessor.from_pretrained(model_path)
caption_model.eval()
if not config.caption_offload:
caption_model = caption_model.to(config.device)
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")
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'
# 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 image_analysis_fn(ci,image):
image = image.convert('RGB')
image_features = ci.image_to_features(image)
top_mediums = ci.mediums.rank(image_features, 5)
top_artists = ci.artists.rank(image_features, 5)
top_movements = ci.movements.rank(image_features, 5)
top_trendings = ci.trendings.rank(image_features, 5)
top_flavors = ci.flavors.rank(image_features, 5)
medium_ranks = {medium: sim for medium, sim in zip(top_mediums, ci.similarities(image_features, top_mediums))}
artist_ranks = {artist: sim for artist, sim in zip(top_artists, ci.similarities(image_features, top_artists))}
movement_ranks = {movement: sim for movement, sim in zip(top_movements, ci.similarities(image_features, top_movements))}
trending_ranks = {trending: sim for trending, sim in zip(top_trendings, ci.similarities(image_features, top_trendings))}
flavor_ranks = {flavor: sim for flavor, sim in zip(top_flavors, ci.similarities(image_features, top_flavors))}
return {
"medium_ranks":medium_ranks,
"artist_ranks":artist_ranks,
"movement_ranks":movement_ranks,
"trending_ranks":trending_ranks,
"flavor_ranks":flavor_ranks
}
def generate_sentences(data):
sentences = []
# Get the length of data
data_length = len(data)
# Use a recursive function to handle variable-length data
def generate_recursive(index, current_sentence, current_score):
# Check if recursion is complete
if index == data_length:
sentences.append({"sentence": current_sentence, "score": current_score})
return
# Get the current level data
current_data = data[index]
# Iterate through the current level data
for phrase in current_data:
sentence = current_sentence + ("," if current_sentence.strip() else "") + phrase
score = current_score + current_data[phrase]
generate_recursive(index + 1, sentence, score)
# Start recursive generation of sentences
generate_recursive(0, "", 0)
# Sort the generated sentences by score in descending order
sentences.sort(key=lambda x: x["score"], reverse=True)
def get_random_elements(elements, num):
return random.sample(elements, num)
ps = get_random_elements(sentences, 5)
ps = [s["sentence"] for s in sorted(ps, key=lambda x: x["score"], reverse=True)]
return ps
def image_to_prompt(ci,image, mode):
ci.config.chunk_size = 2048 if ci.config.clip_model_name == "ViT-L-14/openai" else 1024
ci.config.flavor_intermediate_count = 2048 if ci.config.clip_model_name == "ViT-L-14/openai" else 1024
image = image.convert('RGB')
if mode == 'best':
return ci.interrogate(image)
elif mode == 'classic':
return ci.interrogate_classic(image)
elif mode == 'fast':
return ci.interrogate_fast(image)
elif mode == 'negative':
return ci.interrogate_negative(image)
# image = Image.open(image_path).convert('RGB')
# ci = Interrogator(Config(clip_model_name="ViT-L-14/openai"))
# print(ci.interrogate(image))
class ClipInterrogator:
global _available
available=_available
@classmethod
def INPUT_TYPES(s):
return {"required": {
"image": ("IMAGE",),
"prompt_mode": (['fast','classic','best','negative'],),
"image_analysis": (["off","on"],),
},
# "optional":{
# "output":("CLIPINTERROGATOR", {"multiline": True,"default": "", "dynamicPrompts": False})
# },
}
RETURN_TYPES = ("STRING","STRING",)
RETURN_NAMES = ("prompt","random_samples",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Prompt"
OUTPUT_NODE = True
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,True,)
global ci
ci = None
def run(self,image,prompt_mode,image_analysis):
global ci
prompt_mode=prompt_mode[0]
analysis=image_analysis[0]
prompt_result=[]
analysis_result=[]
# 进度条
pbar = comfy.utils.ProgressBar(len(image)*(2 if analysis=='on' else 1))
if ci==None:
config=Config(
clip_model_name="ViT-L-14/openai",
device="cuda" if torch.cuda.is_available() else "cpu",
download_cache=True,
clip_model_path=cache_path,
cache_path=cache_path
)
config.apply_low_vram_defaults()
caption_model,caption_processor=load_caption_model(caption_model_path,config)
config.caption_model= caption_model
config.caption_processor= caption_processor
ci = Interrogator(config)
# else:
# simple_lama.model.to("cuda" if torch.cuda.is_available() else "cpu")
for i in range(len(image)):
im=image[i]
im=tensor2pil(im)
im=im.convert('RGB')
if analysis=='on':
analysis_res=image_analysis_fn(ci,im)
analysis_result.append( analysis_res )
pbar.update(1)
prompt=image_to_prompt(ci,im,prompt_mode)
pbar.update(1)
prompt_result.append(prompt)
# result.save("inpainted.png")
if ci.config.clip_offload and not ci.clip_offloaded:
ci.clip_model = ci.clip_model.to('cpu')
ci.clip_offloaded = True
if ci.config.caption_offload and not ci.caption_offloaded:
ci.caption_model = ci.caption_model.to('cpu')
ci.caption_offloaded = True
# analysis_result=[]
# items = app.graph.getNodeById(31).widgets[2].value["items"]
random_samples=[]
for r in analysis_result:
random_sample = generate_sentences([r['medium_ranks'], r['artist_ranks'],r['movement_ranks'],r['trending_ranks'],r['flavor_ranks']])
for s in random_sample:
random_samples.append(s)
# print(len(random_samples))
# print('-----')
# print( random_samples)
return {
"ui":{
"prompt": prompt_result,
"analysis":analysis_result,
"random_samples":random_samples
},
"result": (prompt_result,random_samples,)}
+272
View File
@@ -0,0 +1,272 @@
#### Thanks:
# [ComfyUI-CLIPSeg](https://github.com/biegert/ComfyUI-CLIPSeg/tree/main)
from transformers import CLIPSegProcessor, CLIPSegForImageSegmentation
from PIL import Image
import torch
import torchvision.transforms as T
import numpy as np
from torchvision.transforms.functional import to_pil_image
import matplotlib.pyplot as plt
import matplotlib.cm as cm
import cv2
from scipy.ndimage import gaussian_filter
from typing import Optional, Tuple
import warnings,os
warnings.filterwarnings("ignore", category=UserWarning, module="torch")
warnings.filterwarnings("ignore", category=UserWarning, module="safetensors")
import folder_paths
import logging
logger = logging.getLogger('CLIPSeg nodes')
clipseg_model_dir = os.path.join(folder_paths.models_dir, "clipseg")
if not os.path.exists(clipseg_model_dir):
clipseg_model_dir='CIDAS/clipseg-rd64-refined'
"""Helper methods for CLIPSeg nodes"""
# 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 tensor_to_numpy(tensor: torch.Tensor) -> np.ndarray:
"""Convert a tensor to a numpy array and scale its values to 0-255."""
array = tensor.numpy().squeeze()
return (array * 255).astype(np.uint8)
def numpy_to_tensor(array: np.ndarray) -> torch.Tensor:
"""Convert a numpy array to a tensor and scale its values from 0-255 to 0-1."""
array = array.astype(np.float32) / 255.0
return torch.from_numpy(array)[None,]
def apply_colormap(mask: torch.Tensor, colormap) -> np.ndarray:
"""Apply a colormap to a tensor and convert it to a numpy array."""
colored_mask = colormap(mask.numpy())[:, :, :3]
return (colored_mask * 255).astype(np.uint8)
def resize_image(image: np.ndarray, dimensions: Tuple[int, int]) -> np.ndarray:
"""Resize an image to the given dimensions using linear interpolation."""
return cv2.resize(image, dimensions, interpolation=cv2.INTER_LINEAR)
def overlay_image(background: np.ndarray, foreground: np.ndarray, alpha: float) -> np.ndarray:
"""Overlay the foreground image onto the background with a given opacity (alpha)."""
return cv2.addWeighted(background, 1 - alpha, foreground, alpha, 0)
def dilate_mask(mask: torch.Tensor, dilation_factor: float) -> torch.Tensor:
"""Dilate a mask using a square kernel with a given dilation factor."""
kernel_size = int(dilation_factor * 2) + 1
kernel = np.ones((kernel_size, kernel_size), np.uint8)
mask_dilated = cv2.dilate(mask.numpy(), kernel, iterations=1)
return torch.from_numpy(mask_dilated)
class CLIPSeg:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
"""
Return a dictionary which contains config for all input fields.
Some types (string): "MODEL", "VAE", "CLIP", "CONDITIONING", "LATENT", "IMAGE", "INT", "STRING", "FLOAT".
Input types "INT", "STRING" or "FLOAT" are special values for fields on the node.
The type can be a list for selection.
Returns: `dict`:
- Key input_fields_group (`string`): Can be either required, hidden or optional. A node class must have property `required`
- Value input_fields (`dict`): Contains input fields config:
* Key field_name (`string`): Name of a entry-point method's argument
* Value field_config (`tuple`):
+ First value is a string indicate the type of field or a list for selection.
+ Secound value is a config for type "INT", "STRING" or "FLOAT".
"""
return {"required":
{
"image": ("IMAGE",),
"text": ("STRING", {"multiline": False}),
},
"optional":
{
"blur": ("FLOAT", {"min": 0, "max": 15, "step": 0.1, "default": 7}),
"threshold": ("FLOAT", {"min": 0, "max": 1, "step": 0.05, "default": 0.4}),
"dilation_factor": ("INT", {"min": 0, "max": 10, "step": 1, "default": 4}),
}
}
CATEGORY = "♾️Mixlab/mask"
RETURN_TYPES = ("MASK", "IMAGE", "IMAGE",)
RETURN_NAMES = ("Mask","Heatmap Mask", "BW Mask")
# INPUT_IS_LIST = True
OUTPUT_IS_LIST = (False,False,False,)
FUNCTION = "segment_image"
def segment_image(self, image: torch.Tensor, text: str, blur: float, threshold: float, dilation_factor: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Create a segmentation mask from an image and a text prompt using CLIPSeg.
Args:
image (torch.Tensor): The image to segment.
text (str): The text prompt to use for segmentation.
blur (float): How much to blur the segmentation mask.
threshold (float): The threshold to use for binarizing the segmentation mask.
dilation_factor (int): How much to dilate the segmentation mask.
Returns:
Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: The segmentation mask, the heatmap mask, and the binarized mask.
"""
# Convert the Tensor to a PIL image
image_np = image.numpy().squeeze() # Remove the first dimension (batch size of 1)
# Convert the numpy array back to the original range (0-255) and data type (uint8)
image_np = (image_np * 255).astype(np.uint8)
# Create a PIL image from the numpy array
i = Image.fromarray(image_np, mode="RGB")
processor = CLIPSegProcessor.from_pretrained(clipseg_model_dir)
model = CLIPSegForImageSegmentation.from_pretrained(clipseg_model_dir)
prompt = text
input_prc = processor(text=prompt, images=i, padding="max_length", return_tensors="pt")
# Predict the segemntation mask
with torch.no_grad():
outputs = model(**input_prc)
tensor = torch.sigmoid(outputs[0]) # get the mask
# Apply a threshold to the original tensor to cut off low values
thresh = threshold
tensor_thresholded = torch.where(tensor > thresh, tensor, torch.tensor(0, dtype=torch.float))
# Apply Gaussian blur to the thresholded tensor
sigma = blur
tensor_smoothed = gaussian_filter(tensor_thresholded.numpy(), sigma=sigma)
tensor_smoothed = torch.from_numpy(tensor_smoothed)
# Normalize the smoothed tensor to [0, 1]
mask_normalized = (tensor_smoothed - tensor_smoothed.min()) / (tensor_smoothed.max() - tensor_smoothed.min())
# Dilate the normalized mask
mask_dilated = dilate_mask(mask_normalized, dilation_factor)
# Convert the mask to a heatmap and a binary mask
heatmap = apply_colormap(mask_dilated, cm.viridis)
binary_mask = apply_colormap(mask_dilated, cm.Greys_r)
# Overlay the heatmap and binary mask on the original image
dimensions = (image_np.shape[1], image_np.shape[0])
heatmap_resized = resize_image(heatmap, dimensions)
binary_mask_resized = resize_image(binary_mask, dimensions)
alpha_heatmap, alpha_binary = 0.5, 1
overlay_heatmap = overlay_image(image_np, heatmap_resized, alpha_heatmap)
overlay_binary = overlay_image(image_np, binary_mask_resized, alpha_binary)
# Convert the numpy arrays to tensors
image_out_heatmap = numpy_to_tensor(overlay_heatmap)
image_out_binary = numpy_to_tensor(overlay_binary)
# Save or display the resulting binary mask
binary_mask_image = Image.fromarray(binary_mask_resized[..., 0])
# convert PIL image to numpy array
tensor_bw = binary_mask_image.convert("L")
tensor_bw=pil2tensor(tensor_bw)
# tensor_bw = np.array(tensor_bw).astype(np.float32) / 255.0
# tensor_bw = torch.from_numpy(tensor_bw)[None,]
# tensor_bw = tensor_bw.squeeze(0)[..., 0]
return (tensor_bw, image_out_heatmap, image_out_binary,)
#OUTPUT_NODE = False
class CombineMasks:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {"required":
{
"input_image": ("IMAGE", ),
"mask_1": ("MASK", ),
"mask_2": ("MASK", ),
},
"optional":
{
"mask_3": ("MASK",),
},
}
CATEGORY = "♾️Mixlab/mask"
RETURN_TYPES = ("MASK", "IMAGE", "IMAGE",)
RETURN_NAMES = ("Combined Mask","Heatmap Mask", "BW Mask")
FUNCTION = "combine_masks"
def combine_masks(self, input_image: torch.Tensor, mask_1: torch.Tensor, mask_2: torch.Tensor, mask_3: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""A method that combines two or three masks into one mask. Takes in tensors and returns the mask as a tensor, as well as the heatmap and binary mask as tensors."""
# Combine masks
if mask_1 is not None:
mask_1 = mask_1.squeeze()
if mask_2 is not None:
mask_2 = mask_2.squeeze()
if mask_3 is not None:
mask_3 = mask_3.squeeze()
print(mask_1.shape,mask_2.shape , mask_3.shape)
combined_mask = mask_1 + mask_2 + mask_3 if mask_3 is not None else mask_1 + mask_2
# print(combined_mask)
# Convert image and masks to numpy arrays
image_np = tensor_to_numpy(input_image)
heatmap = apply_colormap(combined_mask, cm.viridis)
binary_mask = apply_colormap(combined_mask, cm.Greys_r)
# Resize heatmap and binary mask to match the original image dimensions
dimensions = (image_np.shape[1], image_np.shape[0])
print('heatmap',heatmap)
if dimensions is None or dimensions[0] == 0 or dimensions[1] == 0:
raise ValueError("Invalid dimensions")
heatmap_resized = resize_image(heatmap, dimensions)
binary_mask_resized = resize_image(binary_mask, dimensions)
# Overlay the heatmap and binary mask onto the original image
alpha_heatmap, alpha_binary = 0.5, 1
overlay_heatmap = overlay_image(image_np, heatmap_resized, alpha_heatmap)
overlay_binary = overlay_image(image_np, binary_mask_resized, alpha_binary)
# Convert overlays to tensors
image_out_heatmap = numpy_to_tensor(overlay_heatmap)
image_out_binary = numpy_to_tensor(overlay_binary)
return combined_mask, image_out_heatmap, image_out_binary
# A dictionary that contains all nodes you want to export with their names
# NOTE: names should be globally unique
# NODE_CLASS_MAPPINGS = {
# "CLIPSeg": CLIPSeg,
# "CombineSegMasks": CombineMasks,
# }
View File
+340 -1794
View File
File diff suppressed because it is too large Load Diff
-123
View File
@@ -1,123 +0,0 @@
import os,sys
import folder_paths
from PIL import Image
import importlib.util
import numpy as np
import torch
global _available
_available=False
def is_installed(package):
try:
spec = importlib.util.find_spec(package)
except ModuleNotFoundError:
return False
return spec is not None
if is_installed('simple_lama_inpainting')==False:
import subprocess
from packaging import version
if version.parse(torch.__version__)>=version.parse('2.1'):
# 安装
print('#pip install simple_lama_inpainting')
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'simple_lama_inpainting'], capture_output=True, text=True)
#检查命令执行结果
if result.returncode == 0:
print("#install success")
from simple_lama_inpainting import SimpleLama
_available=True
else:
print("#install error")
else:
print('#pls check your torch version >= 2.1')
else:
from simple_lama_inpainting import SimpleLama
_available=True
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")
else:
os.environ['LAMA_MODEL'] = llma_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)
# simple_lama = SimpleLama()
# img_path = "image.png"
# mask_path = "mask.png"
# image = Image.open(img_path)
# mask = Image.open(mask_path).convert('L')
# result = simple_lama(image, mask)
# result.save("inpainted.png")
class LaMaInpainting:
global _available
available=_available
@classmethod
def INPUT_TYPES(s):
return {"required": {
"image": ("IMAGE",),
"mask": ("MASK",),
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Image"
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,)
global simple_lama
simple_lama = None
def run(self,image,mask):
global simple_lama
result=[]
if simple_lama==None:
simple_lama = SimpleLama()
else:
simple_lama.model.to("cuda" if torch.cuda.is_available() else "cpu")
for i in range(len(image)):
im=image[i]
ma=mask[i]
im=tensor2pil(im)
ma=tensor2pil(ma)
ma =ma.convert('L')
res = simple_lama(im, ma)
res=pil2tensor(res)
result.append(res)
# result.save("inpainted.png")
if simple_lama.device=='cuda':
simple_lama.model.to('cpu')
return (result,)
-274
View File
@@ -1,274 +0,0 @@
import scipy.ndimage
import torch
import numpy as np
# from PIL import Image, ImageDraw
from PIL import Image, ImageOps
from comfy.cli_args import args
import cv2,os
from nodes import MAX_RESOLUTION, SaveImage, common_ksampler
import folder_paths,random
# 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 add_masks(mask1, mask2):
mask1 = mask1.cpu()
mask2 = mask2.cpu()
cv2_mask1 = np.array(mask1) * 255
cv2_mask2 = np.array(mask2) * 255
if cv2_mask1.shape == cv2_mask2.shape:
cv2_mask = cv2.add(cv2_mask1, cv2_mask2)
return torch.clamp(torch.from_numpy(cv2_mask) / 255.0, min=0, max=1)
else:
return mask1
def grow(mask, expand, tapered_corners):
c = 0 if tapered_corners else 1
kernel = np.array([[c, 1, c],
[1, 1, 1],
[c, 1, c]])
mask = mask.reshape((-1, mask.shape[-2], mask.shape[-1]))
out = []
for m in mask:
output = m.numpy()
for _ in range(abs(expand)):
if expand < 0:
output = scipy.ndimage.grey_erosion(output, footprint=kernel)
else:
output = scipy.ndimage.grey_dilation(output, footprint=kernel)
output = torch.from_numpy(output)
out.append(output)
return torch.stack(out, dim=0)
def combine(destination, source, x, y):
output = destination.reshape((-1, destination.shape[-2], destination.shape[-1])).clone()
source = source.reshape((-1, source.shape[-2], source.shape[-1]))
left, top = (x, y,)
right, bottom = (min(left + source.shape[-1], destination.shape[-1]), min(top + source.shape[-2], destination.shape[-2]))
visible_width, visible_height = (right - left, bottom - top,)
source_portion = source[:, :visible_height, :visible_width]
destination_portion = destination[:, top:bottom, left:right]
#operation == "subtract":
output[:, top:bottom, left:right] = destination_portion - source_portion
output = torch.clamp(output, 0.0, 1.0)
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
def INPUT_TYPES(s):
return {
"required": {
"mask": ("MASK",),
"outline_width":("INT", {"default": 10,"min": 1, "max": MAX_RESOLUTION, "step": 1}),
"tapered_corners": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ('MASK',)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Mask"
# 运行的函数
def run(self, mask, outline_width, tapered_corners):
m1=grow(mask,outline_width,tapered_corners)
m2=grow(mask,-outline_width,tapered_corners)
m3=combine(m1,m2,0,0)
return (m3,)
class MaskListReplace:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"masks": ("MASK",),
"mask_replace": ("MASK",),
"start_index":("INT", {"default": 0, "min": 0, "step": 1}),
"end_index":("INT", {"default": 0, "min": 0, "step": 1}),
"invert": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("MASK",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Video"
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,)
def run(self, masks,mask_replace,start_index,end_index,invert):
mask_replace=mask_replace[0]
start_index=start_index[0]
end_index=end_index[0]
invert=invert[0]
new_masks=[]
for i in range(len(masks)):
if i>=start_index and i<=end_index:
if invert:
new_masks.append(masks[i])
else:
new_masks.append(mask_replace)
else:
if invert:
new_masks.append(mask_replace)
else:
new_masks.append(masks[i])
return (new_masks,)
class MaskListMerge:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"masks": ("MASK",),
}
}
RETURN_TYPES = ("MASK",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Mask"
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (False,)
def run(self, masks):
mask=masks[0]
if isinstance(masks, list):
for m in masks:
# print(m.shape)
mask = add_masks(mask, m)
return (mask,)
class FeatheredMask:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mask": ("MASK",),
"start_offset":("INT", {"default": 1,
"min": -150,
"max": 150,
"step": 1,
"display": "slider"}),
"feathering_weight":("FLOAT", {"default": 0.1,
"min": 0.0,
"max": 1,
"step": 0.1,
"display": "slider"})
}
}
RETURN_TYPES = ('MASK',)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Mask"
OUTPUT_IS_LIST = (True,)
# 运行的函数
def run(self,mask,start_offset, feathering_weight):
# print(mask.shape,mask.size())
num,_,_=mask.size()
masks=[]
for i in range(num):
mm=mask[i]
image=tensor2pil(mm)
# Open the image using PIL
image = image.convert("L")
if start_offset>0:
image=ImageOps.invert(image)
# Convert the image to a numpy array
image_np = np.array(image)
# Use Canny edge detection to get black contours
edges = cv2.Canny(image_np, 30, 150)
for i in range(0,abs(start_offset)):
# int(100*feathering_weight)
a=int(abs(start_offset)*0.1*i)
# Dilate the black contours to make them wider
kernel = np.ones((a, a), np.uint8)
dilated_edges = cv2.dilate(edges, kernel, iterations=1)
# dilated_edges = cv2.erode(edges, kernel, iterations=1)
# Smooth the dilated edges using Gaussian blur
smoothed_edges = cv2.GaussianBlur(dilated_edges, (5, 5), 0)
# Adjust the feathering weight
feathering_weight = max(0, min(feathering_weight, 1))
# Blend the smoothed edges with the original image to achieve feathering effect
image_np = cv2.addWeighted(image_np, 1, smoothed_edges, feathering_weight, feathering_weight)
# Convert the result back to PIL image
result_image = Image.fromarray(np.uint8(image_np))
result_image=result_image.convert("L")
if start_offset>0:
result_image=ImageOps.invert(result_image)
result_image=result_image.convert("L")
mt=pil2tensor(result_image)
masks.append(mt)
# print( mt.size())
return (masks,)
+52 -552
View File
@@ -1,91 +1,15 @@
import random
import comfy.utils
import os
import numpy as np
from urllib import request, parse
import folder_paths
from PIL import Image, ImageOps,ImageFilter,ImageEnhance,ImageDraw,ImageSequence, ImageFont
from PIL.PngImagePlugin import PngInfo
import hashlib
import requests
import json
# def queue_prompt(prompt_workflow):
# p = {"prompt": prompt_workflow}
# data = json.dumps(p).encode('utf-8')
# 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_files_with_extension(directory, extension):
file_list = []
for root, dirs, files in os.walk(directory):
for file in files:
if file.endswith(extension):
file_name = os.path.splitext(file)[0]
file_list.append(file_name)
return file_list
def join_with_(text_list,delimiter):
joined_text = delimiter.join(text_list)
return joined_text
from urllib import request, parse
def load_json(file_path):
try:
with open(file_path, 'r') as json_file:
data = json.load(json_file)
return data
except FileNotFoundError:
print(f"File not found: {file_path}")
return None
except json.JSONDecodeError:
print(f"Error decoding JSON in file: {file_path}")
return None
def save_json(data_dict, file_path):
try:
with open(file_path, 'w') as json_file:
json.dump(data_dict, json_file, indent=4)
print(f"Data saved to {file_path}")
except Exception as e:
print(f"Error saving JSON to file: {e}")
# pysss的lora加载器
# def get_model_version_info(hash_value):
# # http://127.0.0.1:1082
# proxies = {'http': 'http://127.0.0.1:1082', 'https': 'https://127.0.0.1:1082'}
# api_url = f"https://civitai.com/api/v1/model-versions/by-hash/{hash_value}"
# print(api_url)
# response = requests.get(api_url,proxies=proxies, verify=False)
# if response.status_code == 200:
# return response.json()
# else:
# return None
# def calculate_sha256(file_path):
# sha256_hash = hashlib.sha256()
# with open(file_path, "rb") as f:
# for chunk in iter(lambda: f.read(4096), b""):
# sha256_hash.update(chunk)
# return sha256_hash.hexdigest()
class AnyType(str):
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
def __ne__(self, __value: object) -> bool:
return False
any_type = AnyType("*")
def queue_prompt(prompt_workflow):
p = {"prompt": prompt_workflow}
data = json.dumps(p).encode('utf-8')
req = request.Request("http://127.0.0.1:8188/prompt", data=data)
request.urlopen(req)
default_prompt1='''Swing
@@ -121,232 +45,6 @@ default_prompt1='''Swing
'''
default_prompt1="\n".join([p.strip() for p in default_prompt1.split('\n') if p.strip()!=''])
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
def addWeight(text, weight=1):
if weight == 1:
return text
else:
return f"({text}:{round(weight,3)})"
def prompt_delete_words(sentence, new_words_length):
# 使用逗号分割句子,并去除空格
words = [word.strip() for word in sentence.split(",")]
# 计算需要删除的单词数量
num_to_delete = len(words) - new_words_length
words_to=[w for w in words]
# 逐个删除单词并存储在新列表中
new_words = []
for i in range(len(words)):
if num_to_delete > 0:
num_to_delete -= 1
else:
words_to.pop()
if len(words_to)>0:
new_words.append(", ".join(words_to))
return new_words
# # 测试方法
# sentence = "a computer, a glass tablet with a keyboard on a dark background, 3d illustration, reflection, cgi 8k, clear glass, archaic, cut-away, white outline"
# new_words_length = 5
# result = prompt_delete_words(sentence, new_words_length)
# print(result)
class PromptImage:
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
self.type = "output"
self.prefix_append = "PromptImage"
self.compress_level = 4
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"prompts": ("STRING",
{
"multiline": True,
"default": '',
"dynamicPrompts": False
}),
"images": ("IMAGE",{"default": None}),
"save_to_image": (["enable", "disable"],),
}
}
RETURN_TYPES = ()
OUTPUT_NODE = True
INPUT_IS_LIST = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Output"
# 运行的函数
def run(self,prompts,images,save_to_image):
filename_prefix="mixlab_"
filename_prefix += self.prefix_append
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
filename_prefix, self.output_dir, images[0].shape[1], images[0].shape[0])
results = list()
save_to_image=save_to_image[0]=='enable'
for index in range(len(images)):
res=[]
imgs=images[index]
for image in imgs:
img=tensor2pil(image)
metadata = None
if save_to_image:
metadata = PngInfo()
prompt_text=prompts[index]
if prompt_text is not None:
metadata.add_text("prompt_text", prompt_text)
file = f"{filename}_{index}_{counter:05}_.png"
img.save(os.path.join(full_output_folder, file), pnginfo=metadata, compress_level=self.compress_level)
res.append({
"filename": file,
"subfolder": subfolder,
"type": self.type
})
counter += 1
results.append(res)
return { "ui": { "_images": results,"prompts":prompts } }
class PromptSimplification:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"prompt": ("STRING",
{
"multiline": True,
"default": '',
"dynamicPrompts": False
}),
"length":("INT", {"default": 5, "min": 1,"max":100, "step": 1, "display": "number"}),
# "min_value":("FLOAT", {
# "default": -2,
# "min": -10,
# "max": 0xffffffffffffffff,
# "step": 0.01,
# "display": "number"
# }),
# "max_value":("FLOAT", {
# "default": 2,
# "min": -10,
# "max": 0xffffffffffffffff,
# "step": 0.01,
# "display": "number"
# }),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("prompts",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Prompt"
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,)
OUTPUT_NODE = True
# 运行的函数
def run(self,prompt,length):
length=length[0]
result=[]
for p in prompt:
nps=prompt_delete_words(p,length)
for n in nps:
result.append(n)
result= [elem.strip() for elem in result if elem.strip()]
return {"ui": {"prompts": result}, "result": (result,)}
class PromptSlide:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"prompt_keyword": ("STRING",
{
"multiline": False,
"default": '',
"dynamicPrompts": False
}),
"weight":("FLOAT", {"default": 1, "min": -3,"max": 3,
"step": 0.01,
"display": "slider"}),
# "min_value":("FLOAT", {
# "default": -2,
# "min": -10,
# "max": 0xffffffffffffffff,
# "step": 0.01,
# "display": "number"
# }),
# "max_value":("FLOAT", {
# "default": 2,
# "min": -10,
# "max": 0xffffffffffffffff,
# "step": 0.01,
# "display": "number"
# }),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("prompt",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Prompt"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
OUTPUT_NODE = False
# 运行的函数
def run(self,prompt_keyword,weight):
# if weight < min_value:
# weight= min_value
# elif weight > max_value:
# weight= max_value
p=addWeight(prompt_keyword,weight)
return (p,)
class RandomPrompt:
'''
@@ -372,10 +70,6 @@ class RandomPrompt:
"default": 'sticker, Cartoon, ``'
}),
"random_sample": (["enable", "disable"],),
# "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
},
"optional":{
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
@@ -385,15 +79,15 @@ class RandomPrompt:
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Prompt"
CATEGORY = "♾️Mixlab/prompt"
OUTPUT_IS_LIST = (True,)
OUTPUT_NODE = True
# 运行的函数
def run(self,max_count,mutable_prompt,immutable_prompt,random_sample,seed=0):
# print('#运行的函数',mutable_prompt,immutable_prompt,max_count,random_sample)
def run(self,max_count,mutable_prompt,immutable_prompt,random_sample):
print('#运行的函数',mutable_prompt,immutable_prompt,max_count,random_sample)
# Split the text into an array of words
words1 = mutable_prompt.split("\n")
@@ -412,11 +106,6 @@ class RandomPrompt:
w1=w1.strip()
for w2 in words2:
w2=w2.strip()
if '``' not in w2:
if w2=="":
w2='``'
else:
w2=w2+',``'
if w1!='' and w2!='':
prompts.append(w2.replace('``', w1))
pbar.update(1)
@@ -430,258 +119,69 @@ class RandomPrompt:
else:
prompts = prompts[:min(max_count,len(prompts))]
prompts= [elem.strip() for elem in prompts if elem.strip()]
# return (new_prompt)
return {"ui": {"prompts": prompts}, "result": (prompts,)}
# class LoraPrompt:
# @classmethod
# def INPUT_TYPES(s):
# return {
# "required": {
# "lora_name":(sorted(folder_paths.get_filename_list("loras"), key=str.lower),),
# "weight": ("FLOAT", {"default": 1, "min": -2, "max": 2,"step":0.01 ,"display": "slider"}),
# "force_update": ("BOOLEAN", {"default": False}),
# },
# }
# RETURN_TYPES = ("STRING","STRING",any_type)
# RETURN_NAMES = ("lora_name","prompt","tags",)
# FUNCTION = "run"
# CATEGORY = "♾️Mixlab/Prompt"
# OUTPUT_IS_LIST = (False,False,True,)
# # OUTPUT_NODE = True
# # 运行的函数
# def run(self,lora_name,weight,force_update=False):
# # print('##LoraPrompt',__file__)
# # 从本地数据库读取
# json_tags_path = os.path.join(os.path.dirname(os.path.dirname(__file__)),r'data/loras_tags.json')
# if not os.path.exists(json_tags_path):
# save_json({},json_tags_path)
# lora_tags = load_json(json_tags_path)
# output_tags = lora_tags.get(lora_name, None) if lora_tags is not None else None
# if output_tags is not None:
# output_tags = ",".join(output_tags)
# print("trainedWords:",output_tags)
# else:
# output_tags = ""
# lora_path = folder_paths.get_full_path("loras", lora_name)
# if output_tags == "" or force_update:
# print("calculating lora hash")
# LORAsha256 = calculate_sha256(lora_path)
# print("requesting infos")
# model_info = get_model_version_info(LORAsha256)
# if model_info is not None:
# if "trainedWords" in model_info:
# print("tags found!")
# if lora_tags is None:
# lora_tags = {}
# lora_tags[lora_name] = model_info["trainedWords"]
# save_json(lora_tags,json_tags_path)
# output_tags = ",".join(model_info["trainedWords"])
# print("trainedWords:",output_tags)
# else:
# print("No informations found.")
# if lora_tags is None:
# lora_tags = {}
# lora_tags[lora_name] = []
# save_json(lora_tags,json_tags_path)
# weight = round(weight, 3)
# prompt=[]
# for p in output_tags.split(','):
# if weight!=1:
# prompt.append('('+p+':'+str(weight)+')')
# else:
# prompt.append(p)
# prompt=",".join(prompt)
# return (lora_name,prompt,output_tags.split(','),)
class EmbeddingPrompt:
class RunWorkflow:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"embedding":(folder_paths.get_filename_list("embeddings"),),
"weight": ("FLOAT", {"default": 1, "min": -2, "max": 2,"step":0.01 ,"display": "slider"}),
"workflow": ("STRING", {
"multiline": False,
"default": ''
}),
"prompt": ("STRING", {
"multiline": False,
"default": ''
}),
"image": ("IMAGE",),
"input_node": ("STRING", {
"multiline": False,
"default": ''
}),
"output_node": ("STRING", {
"multiline": False,
"default": ''
}),
},
}
RETURN_TYPES = ("STRING",)
RETURN_TYPES = ("IMAGE","STRING",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Prompt"
CATEGORY = "♾️Mixlab/workflow"
OUTPUT_IS_LIST = (True,)
OUTPUT_NODE = True
OUTPUT_IS_LIST = (False,)
# OUTPUT_NODE = True
# 运行的函数
def run(self,embedding,weight):
weight = round(weight, 3)
prompt='embedding:'+embedding
if weight!=1:
prompt='('+prompt+':'+str(weight)+')'
prompt=" "+prompt+' '
def run(self,workflow,prompt,image,input_node,output_node):
print('#运行的函数',prompt,image,input_node,output_node)
workflow=json.loads(workflow)
input_node=input_node.split(".")
workflow[input_node[0]][input_node[1]][input_node[2]]=prompt
workflow_new={}
# 遍历,seed设为随机
for key, value in workflow.items():
if 'inputs' in value:
if 'seed' in value['inputs']:
value['inputs']['seed']= random.randint(1, 18446744073709551614)
workflow_new[key]=value
queue_prompt(workflow_new)
print('#运行的函数',workflow_new[input_node[0]])
# return (new_prompt)
return (prompt,)
return {"ui":{"images": []},"result": ([image],['text'],)}
# RETURN_TYPES = (any_type,)
# conditioning :提示,正向or负向
# clip:clip模型
# gligen_textbox_model:gligen模型
# grids:矩形框的集合
# labels:每个矩形框对应的标签的集合
# index:选取第几个矩形框作为gligen的box
class GLIGENTextBoxApply_Advanced:
@classmethod
def INPUT_TYPES(s):
return {"required": {"conditioning": ("CONDITIONING", ),
"clip": ("CLIP", ),
"gligen_textbox_model": ("GLIGEN", ),
"grids": ("_GRID",),
"labels": ("STRING",
{
"multiline": True,
"default": "",
"forceInput": True
}),
"index": ("INT", {"default": -1, "min": -1, "max": 300, "step": 1}),
"max_size": ("INT", {"default": 8, "min": 1, "max": 300, "step": 1}),
"random_shuffle":(["on","off"],),
},
"optional":{
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff,"step": 1}),
}
}
RETURN_TYPES = ("CONDITIONING","STRING",)
RETURN_NAMES = ("CONDITIONING","label",)
FUNCTION = "run"
# INPUT_IS_LIST = True
CATEGORY = "♾️Mixlab/Prompt"
def run(self, conditioning, clip, gligen_textbox_model, grids, labels, index,max_size,random_shuffle,seed=0):
# print('grids',grids)
# conditioning=conditioning[0]
# clip=clip[0]
# gligen_textbox_model=gligen_textbox_model[0]
# index=index[0]
# max_size=max_size[0]
# random_shuffle=random_shuffle[0]
texts=labels
if index>-1:
texts=[labels[index]]
grids=[grids[index]]
if random_shuffle=='on':
sss=[[texts[i],grids[i]] for i in range(len(texts))]
random.shuffle(sss)
texts=[s[0] for s in sss]
grids=[s[1] for s in sss]
if len(texts) > max_size:
texts = texts[:max_size]
c = []
for t in conditioning:
n = [t[0], t[1].copy()]
# 多个
position_params=[]
for i in range(len(texts)):
text=texts[i]
grid=grids[i]
x,y,width,height=grid
# print(text)
cond, cond_pooled = clip.encode_from_tokens(clip.tokenize(text), return_pooled=True)
position_params =position_params+ [(cond_pooled, height // 8, width // 8, y // 8, x // 8)]
# 前一个
prev = []
if "gligen" in n[1]:
prev = n[1]['gligen'][2]
n[1]['gligen'] = ("position", gligen_textbox_model, prev + position_params)
# print('gligen',n)
c.append(n)
# 下面这个写法有bug
# for i in range(len(texts)):
# text=texts[i]
# grid=grids[i]
# x,y,width,height=grid
# cond, cond_pooled = clip.encode_from_tokens(clip.tokenize(text), return_pooled=True)
# for t in conditioning:
# n = [t[0], t[1].copy()]
# position_params = [(cond_pooled, height // 8, width // 8, y // 8, x // 8)]
# prev = []
# if "gligen" in n[1]:
# prev = n[1]['gligen'][2]
# n[1]['gligen'] = ("position", gligen_textbox_model, prev + position_params)
# c.append(n)
return (c,texts, )
class JoinWithDelimiter:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"text_list": (any_type,),
"delimiter":(["newline","comma","backslash","space"],),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Text"
INPUT_IS_LIST = True # 当true的时候,输入时list,当false的时候,如果输入是list,则会自动包一层for循环调用
OUTPUT_IS_LIST = (False,)
def run(self,text_list,delimiter):
delimiter=delimiter[0]
if delimiter =='newline':
delimiter='\n'
elif delimiter=='comma':
delimiter=','
elif delimiter=='backslash':
delimiter='\\'
elif delimiter=='space':
delimiter=' '
t=''
if isinstance(text_list, list):
t=join_with_(text_list,delimiter)
return (t,)
-708
View File
@@ -1,708 +0,0 @@
import os,sys
import folder_paths
from PIL import Image
import importlib.util
import comfy.utils
import numpy as np
import torch
from huggingface_hub import hf_hub_download
import torch.nn as nn
import torch.nn.functional as F
from torchvision.transforms.functional import normalize
# BRIA-RMBG-1.4 / briarmbg.py
class REBNCONV(nn.Module):
def __init__(self,in_ch=3,out_ch=3,dirate=1,stride=1):
super(REBNCONV,self).__init__()
self.conv_s1 = nn.Conv2d(in_ch,out_ch,3,padding=1*dirate,dilation=1*dirate,stride=stride)
self.bn_s1 = nn.BatchNorm2d(out_ch)
self.relu_s1 = nn.ReLU(inplace=True)
def forward(self,x):
hx = x
xout = self.relu_s1(self.bn_s1(self.conv_s1(hx)))
return xout
## upsample tensor 'src' to have the same spatial size with tensor 'tar'
def _upsample_like(src,tar):
src = F.interpolate(src,size=tar.shape[2:],mode='bilinear')
return src
### RSU-7 ###
class RSU7(nn.Module):
def __init__(self, in_ch=3, mid_ch=12, out_ch=3, img_size=512):
super(RSU7,self).__init__()
self.in_ch = in_ch
self.mid_ch = mid_ch
self.out_ch = out_ch
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1) ## 1 -> 1/2
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool5 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.rebnconv7 = REBNCONV(mid_ch,mid_ch,dirate=2)
self.rebnconv6d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
def forward(self,x):
b, c, h, w = x.shape
hx = x
hxin = self.rebnconvin(hx)
hx1 = self.rebnconv1(hxin)
hx = self.pool1(hx1)
hx2 = self.rebnconv2(hx)
hx = self.pool2(hx2)
hx3 = self.rebnconv3(hx)
hx = self.pool3(hx3)
hx4 = self.rebnconv4(hx)
hx = self.pool4(hx4)
hx5 = self.rebnconv5(hx)
hx = self.pool5(hx5)
hx6 = self.rebnconv6(hx)
hx7 = self.rebnconv7(hx6)
hx6d = self.rebnconv6d(torch.cat((hx7,hx6),1))
hx6dup = _upsample_like(hx6d,hx5)
hx5d = self.rebnconv5d(torch.cat((hx6dup,hx5),1))
hx5dup = _upsample_like(hx5d,hx4)
hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))
hx4dup = _upsample_like(hx4d,hx3)
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
hx3dup = _upsample_like(hx3d,hx2)
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
hx2dup = _upsample_like(hx2d,hx1)
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
return hx1d + hxin
### RSU-6 ###
class RSU6(nn.Module):
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
super(RSU6,self).__init__()
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=2)
self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
def forward(self,x):
hx = x
hxin = self.rebnconvin(hx)
hx1 = self.rebnconv1(hxin)
hx = self.pool1(hx1)
hx2 = self.rebnconv2(hx)
hx = self.pool2(hx2)
hx3 = self.rebnconv3(hx)
hx = self.pool3(hx3)
hx4 = self.rebnconv4(hx)
hx = self.pool4(hx4)
hx5 = self.rebnconv5(hx)
hx6 = self.rebnconv6(hx5)
hx5d = self.rebnconv5d(torch.cat((hx6,hx5),1))
hx5dup = _upsample_like(hx5d,hx4)
hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))
hx4dup = _upsample_like(hx4d,hx3)
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
hx3dup = _upsample_like(hx3d,hx2)
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
hx2dup = _upsample_like(hx2d,hx1)
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
return hx1d + hxin
### RSU-5 ###
class RSU5(nn.Module):
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
super(RSU5,self).__init__()
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=2)
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
def forward(self,x):
hx = x
hxin = self.rebnconvin(hx)
hx1 = self.rebnconv1(hxin)
hx = self.pool1(hx1)
hx2 = self.rebnconv2(hx)
hx = self.pool2(hx2)
hx3 = self.rebnconv3(hx)
hx = self.pool3(hx3)
hx4 = self.rebnconv4(hx)
hx5 = self.rebnconv5(hx4)
hx4d = self.rebnconv4d(torch.cat((hx5,hx4),1))
hx4dup = _upsample_like(hx4d,hx3)
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
hx3dup = _upsample_like(hx3d,hx2)
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
hx2dup = _upsample_like(hx2d,hx1)
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
return hx1d + hxin
### RSU-4 ###
class RSU4(nn.Module):
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
super(RSU4,self).__init__()
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=2)
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
def forward(self,x):
hx = x
hxin = self.rebnconvin(hx)
hx1 = self.rebnconv1(hxin)
hx = self.pool1(hx1)
hx2 = self.rebnconv2(hx)
hx = self.pool2(hx2)
hx3 = self.rebnconv3(hx)
hx4 = self.rebnconv4(hx3)
hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))
hx3dup = _upsample_like(hx3d,hx2)
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
hx2dup = _upsample_like(hx2d,hx1)
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
return hx1d + hxin
### RSU-4F ###
class RSU4F(nn.Module):
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
super(RSU4F,self).__init__()
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=2)
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=4)
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=8)
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=4)
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=2)
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
def forward(self,x):
hx = x
hxin = self.rebnconvin(hx)
hx1 = self.rebnconv1(hxin)
hx2 = self.rebnconv2(hx1)
hx3 = self.rebnconv3(hx2)
hx4 = self.rebnconv4(hx3)
hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))
hx2d = self.rebnconv2d(torch.cat((hx3d,hx2),1))
hx1d = self.rebnconv1d(torch.cat((hx2d,hx1),1))
return hx1d + hxin
class myrebnconv(nn.Module):
def __init__(self, in_ch=3,
out_ch=1,
kernel_size=3,
stride=1,
padding=1,
dilation=1,
groups=1):
super(myrebnconv,self).__init__()
self.conv = nn.Conv2d(in_ch,
out_ch,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
groups=groups)
self.bn = nn.BatchNorm2d(out_ch)
self.rl = nn.ReLU(inplace=True)
def forward(self,x):
return self.rl(self.bn(self.conv(x)))
class BriaRMBG(nn.Module):
def __init__(self,in_ch=3,out_ch=1):
super(BriaRMBG,self).__init__()
self.conv_in = nn.Conv2d(in_ch,64,3,stride=2,padding=1)
self.pool_in = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage1 = RSU7(64,32,64)
self.pool12 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage2 = RSU6(64,32,128)
self.pool23 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage3 = RSU5(128,64,256)
self.pool34 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage4 = RSU4(256,128,512)
self.pool45 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage5 = RSU4F(512,256,512)
self.pool56 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage6 = RSU4F(512,256,512)
# decoder
self.stage5d = RSU4F(1024,256,512)
self.stage4d = RSU4(1024,128,256)
self.stage3d = RSU5(512,64,128)
self.stage2d = RSU6(256,32,64)
self.stage1d = RSU7(128,16,64)
self.side1 = nn.Conv2d(64,out_ch,3,padding=1)
self.side2 = nn.Conv2d(64,out_ch,3,padding=1)
self.side3 = nn.Conv2d(128,out_ch,3,padding=1)
self.side4 = nn.Conv2d(256,out_ch,3,padding=1)
self.side5 = nn.Conv2d(512,out_ch,3,padding=1)
self.side6 = nn.Conv2d(512,out_ch,3,padding=1)
# self.outconv = nn.Conv2d(6*out_ch,out_ch,1)
def forward(self,x):
hx = x
hxin = self.conv_in(hx)
#hx = self.pool_in(hxin)
#stage 1
hx1 = self.stage1(hxin)
hx = self.pool12(hx1)
#stage 2
hx2 = self.stage2(hx)
hx = self.pool23(hx2)
#stage 3
hx3 = self.stage3(hx)
hx = self.pool34(hx3)
#stage 4
hx4 = self.stage4(hx)
hx = self.pool45(hx4)
#stage 5
hx5 = self.stage5(hx)
hx = self.pool56(hx5)
#stage 6
hx6 = self.stage6(hx)
hx6up = _upsample_like(hx6,hx5)
#-------------------- decoder --------------------
hx5d = self.stage5d(torch.cat((hx6up,hx5),1))
hx5dup = _upsample_like(hx5d,hx4)
hx4d = self.stage4d(torch.cat((hx5dup,hx4),1))
hx4dup = _upsample_like(hx4d,hx3)
hx3d = self.stage3d(torch.cat((hx4dup,hx3),1))
hx3dup = _upsample_like(hx3d,hx2)
hx2d = self.stage2d(torch.cat((hx3dup,hx2),1))
hx2dup = _upsample_like(hx2d,hx1)
hx1d = self.stage1d(torch.cat((hx2dup,hx1),1))
#side output
d1 = self.side1(hx1d)
d1 = _upsample_like(d1,x)
d2 = self.side2(hx2d)
d2 = _upsample_like(d2,x)
d3 = self.side3(hx3d)
d3 = _upsample_like(d3,x)
d4 = self.side4(hx4d)
d4 = _upsample_like(d4,x)
d5 = self.side5(hx5d)
d5 = _upsample_like(d5,x)
d6 = self.side6(hx6)
d6 = _upsample_like(d6,x)
return [F.sigmoid(d1), F.sigmoid(d2), F.sigmoid(d3), F.sigmoid(d4), F.sigmoid(d5), F.sigmoid(d6)],[hx1d,hx2d,hx3d,hx4d,hx5d,hx6]
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=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)
except ModuleNotFoundError:
return False
return spec is not None
try:
if is_installed('rembg')==False:
import subprocess
# 安装
print('#pip install rembg[gpu]')
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'rembg[gpu]'], capture_output=True, text=True)
#检查命令执行结果
if result.returncode == 0:
print("#install success")
from rembg import new_session, remove
_available=True
else:
print("#install error")
else:
from rembg import new_session, remove
_available=True
except:
_available=False
def run_briarmbg(images=[]):
mroot=U2NET_HOME
m=os.path.join(mroot,'briarmbg.pth')
if os.path.exists(m)==False:
# 下载
m1=hf_hub_download("briaai/RMBG-1.4",
local_dir=mroot,
filename='model.pth',
local_dir_use_symlinks=False,
endpoint='https://hf-mirror.com')
os.rename(m1, m)
net=BriaRMBG()
if torch.cuda.is_available():
net.load_state_dict(torch.load(m))
net=net.cuda()
else:
net.load_state_dict(torch.load(m,map_location="cpu"))
net.eval()
masks=[]
rgba_images=[]
rgb_images=[]
for orig_image in images:
w,h = orig_im_size = orig_image.size
image = orig_image.convert('RGB')
model_input_size = (1024, 1024)
image = image.resize(model_input_size, Image.BILINEAR)
im_np = np.array(image)
im_tensor = torch.tensor(im_np, dtype=torch.float32).permute(2,0,1)
im_tensor = torch.unsqueeze(im_tensor,0)
im_tensor = torch.divide(im_tensor,255.0)
im_tensor = normalize(im_tensor,[0.5,0.5,0.5],[1.0,1.0,1.0])
if torch.cuda.is_available():
im_tensor=im_tensor.cuda()
result=net(im_tensor)
result = torch.squeeze(F.interpolate(result[0][0], size=(h,w), mode='bilinear') ,0)
ma = torch.max(result)
mi = torch.min(result)
result = (result-mi)/(ma-mi)
im_array = (result*255).cpu().data.numpy().astype(np.uint8)
mask = Image.fromarray(np.squeeze(im_array))
# mask.save('test.png')
# mask=tensor2pil(result)
mask=mask.convert('L')
masks.append(mask)
# rgba图
image_rgba =orig_image.convert("RGBA")
image_rgba.putalpha(mask)
rgba_images.append(image_rgba)
#rgb
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)
return (masks,rgba_images,rgb_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=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)
# mask=mask.convert('L')
# masks.append(mask)
if model_name=="u2net_cloth_seg":
width, original_height = mask.size
num_slices = original_height // img.height
for i in range(num_slices):
top = i * img.height
bottom = (i + 1) * img.height
slice_image = mask.crop((0, top, width, bottom))
slice_mask=slice_image.convert('L')
masks.append(slice_mask)
# rgba图
image_rgba = img.convert("RGBA")
image_rgba.putalpha(slice_mask)
rgba_images.append(image_rgba)
#rgb
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)
else:
mask=mask.convert('L')
# mask.save(output_path)
masks.append(mask)
# rgba图
image_rgba = img.convert("RGBA")
image_rgba.putalpha(mask)
rgba_images.append(image_rgba)
#rgb
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)
if pbar:
pbar.update(1)
return (masks,rgba_images,rgb_images)
# 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)
class RembgNode_:
global _available
available=_available
@classmethod
def INPUT_TYPES(s):
return {"required": {
"image": ("IMAGE",),
"model_name": (get_rembg_models(U2NET_HOME),),
},
}
RETURN_TYPES = ("MASK","IMAGE","RGBA",)
RETURN_NAMES = ("masks","images","RGBAs")
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Mask"
OUTPUT_NODE = True
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,True,True,)
def run(self,image,model_name):
# 兼容list输入和batch输入
model_name=model_name[0]
images=[]
for ims in image:
for im in ims:
im=tensor2pil(im)
images.append(im)
if model_name=='briarmbg':
masks,rgba_images,rgb_images=run_briarmbg(images)
else:
masks,rgba_images,rgb_images=run_rembg(model_name,images, comfy.utils.ProgressBar(len(images) ))
masks=[pil2tensor(m) for m in masks]
rgba_images=[pil2tensor(m) for m in rgba_images]
rgb_images=[pil2tensor(m) for m in rgb_images]
return (masks,rgb_images,rgba_images,)
+8 -8
View File
@@ -90,10 +90,10 @@ class ScreenShareNode:
} }
RETURN_TYPES = ('IMAGE','STRING','FLOAT',"INT")
RETURN_NAMES = ("current frame (image)","prompt","denoise (float)","seed (int)")
RETURN_NAMES = ("IMAGE","PROMPT","FLOAT","INT")
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Screen"
CATEGORY = "♾️Mixlab/image"
# INPUT_IS_LIST = True
OUTPUT_IS_LIST = (False,False,False,False)
@@ -109,7 +109,7 @@ class FloatingVideo:
@classmethod
def INPUT_TYPES(s):
return { "required":{
"image": ("IMAGE",)
"images": ("IMAGE",)
}, }
# RETURN_TYPES = ('IMAGE','MASK')
@@ -118,22 +118,22 @@ class FloatingVideo:
OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Screen"
CATEGORY = "♾️Mixlab/image"
# INPUT_IS_LIST = True
# OUTPUT_IS_LIST = (False,False,)
# 运行的函数
def run(self,image):
def run(self,images):
results = list()
for im in image:
im=tensor2pil(im)
for image in images:
image=tensor2pil(image)
# image_base64 = base64.b64encode(image.tobytes())
buffered = BytesIO()
im.save(buffered, format="JPEG")
image.save(buffered, format="JPEG")
image_base64 = base64.b64encode(buffered.getvalue()).decode("utf-8")
results.append(image_base64)
+37
View File
@@ -0,0 +1,37 @@
import urllib.parse
# 分享到微博
class ShareToWeibo:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"title":("STRING",{"multiline": True,"default": "","dynamicPrompts": False}),
"pic_url":("STRING",{"multiline": False,"default": "","dynamicPrompts": False}),
"url":("STRING",{"multiline": False,"default": "","dynamicPrompts": False}),
}
}
RETURN_TYPES = ()
# RETURN_NAMES = ("number",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/share"
INPUT_IS_LIST = False
OUTPUT_NODE = True
# OUTPUT_IS_LIST = ()
def run(self, title, pic_url, url):
encoded_title = urllib.parse.quote(title)
encoded_pic_url = urllib.parse.quote(pic_url)
encoded_url = urllib.parse.quote(url)
url = "https://service.weibo.com/share/share.php?title={}&pic={}&url={}".format(encoded_title,encoded_pic_url,encoded_url)
print(url)
return {"ui": {"url": [url]}, "result": ()}
-503
View File
@@ -1,503 +0,0 @@
import comfy
import torch
from dataclasses import dataclass
import torch.nn as nn
from comfy.model_patcher import ModelPatcher
import comfy.ops
from typing import Union
import comfy.sample
import latent_preview
import comfy.utils
T = torch.Tensor
from .VisualStylePrompting.attention_functions import VisualStyleProcessor
class ApplyVisualStylePrompting:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"reference_image": ("IMAGE",),
"reference_image_text": ("STRING", {"multiline": True}),
"model": ("MODEL",),
"clip": ("CLIP", ),
"vae": ("VAE", ),
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING", ),
"enabled": ("BOOLEAN", {"default": True}),
"denoise": ("FLOAT", {"default": 1., "min": 0., "max": 1., "step": 1e-2}),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 4096,"step":2})
}
}
RETURN_TYPES = ("MODEL", "CONDITIONING","CONDITIONING", "LATENT")
RETURN_NAMES = ("model", "positive", "negative", "latents")
CATEGORY = "♾️Mixlab/Style"
FUNCTION = "run"
def run(
self,
reference_image,
reference_image_text,
model: comfy.model_patcher.ModelPatcher,
clip,
vae,
positive,
negative,
enabled,
denoise,
batch_size=1
):
tokens = clip.tokenize(reference_image_text)
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
reference_image_prompt=[[cond, {"pooled_output": pooled}]]
reference_image = reference_image.repeat(((batch_size+1)//2, 1,1,1))
self.model = model
reference_latent = vae.encode(reference_image[:,:,:,:3])
for n, m in model.model.diffusion_model.named_modules():
if m.__class__.__name__ == "CrossAttention":
processor = VisualStyleProcessor(m, enabled=enabled)
setattr(m, 'forward', processor.visual_style_forward)
conditioning_prompt = reference_image_prompt + positive
negative_prompt = negative * 2
latents = torch.zeros_like(reference_latent)
latents = torch.cat([latents] * 2)
if denoise < 1.0:
latents[::1] = reference_latent[:1]
else:
latents[::2] = reference_latent
denoise_mask = torch.ones_like(latents)[:, :1, ...] * denoise
denoise_mask[0] = 0.
return (model, conditioning_prompt, negative_prompt, {"samples": latents, "noise_mask": denoise_mask})
def exists(val):
return val is not None
def default(val, d):
if exists(val):
return val
return d
class StyleAlignedArgs:
def __init__(self, share_attn: str) -> None:
self.adain_keys = "k" in share_attn
self.adain_values = "v" in share_attn
self.adain_queries = "q" in share_attn
share_attention: bool = True
adain_queries: bool = True
adain_keys: bool = True
adain_values: bool = True
def expand_first(
feat: T,
scale=1.0,
) -> T:
"""
Expand the first element so it has the same shape as the rest of the batch.
"""
b = feat.shape[0]
feat_style = torch.stack((feat[0], feat[b // 2])).unsqueeze(1)
if scale == 1:
feat_style = feat_style.expand(2, b // 2, *feat.shape[1:])
else:
feat_style = feat_style.repeat(1, b // 2, 1, 1, 1)
feat_style = torch.cat([feat_style[:, :1], scale * feat_style[:, 1:]], dim=1)
return feat_style.reshape(*feat.shape)
def concat_first(feat: T, dim=2, scale=1.0) -> T:
"""
concat the the feature and the style feature expanded above
"""
feat_style = expand_first(feat, scale=scale)
return torch.cat((feat, feat_style), dim=dim)
def calc_mean_std(feat, eps: float = 1e-5) -> "tuple[T, T]":
feat_std = (feat.var(dim=-2, keepdims=True) + eps).sqrt()
feat_mean = feat.mean(dim=-2, keepdims=True)
return feat_mean, feat_std
def adain(feat: T) -> T:
feat_mean, feat_std = calc_mean_std(feat)
feat_style_mean = expand_first(feat_mean)
feat_style_std = expand_first(feat_std)
feat = (feat - feat_mean) / feat_std
feat = feat * feat_style_std + feat_style_mean
return feat
class SharedAttentionProcessor:
def __init__(self, args: StyleAlignedArgs, scale: float):
self.args = args
self.scale = scale
def __call__(self, q, k, v, extra_options):
if self.args.adain_queries:
q = adain(q)
if self.args.adain_keys:
k = adain(k)
if self.args.adain_values:
v = adain(v)
if self.args.share_attention:
k = concat_first(k, -2, scale=self.scale)
v = concat_first(v, -2)
return q, k, v
def get_norm_layers(
layer: nn.Module,
norm_layers_: "dict[str, list[Union[nn.GroupNorm, nn.LayerNorm]]]",
share_layer_norm: bool,
share_group_norm: bool,
):
if isinstance(layer, nn.LayerNorm) and share_layer_norm:
norm_layers_["layer"].append(layer)
if isinstance(layer, nn.GroupNorm) and share_group_norm:
norm_layers_["group"].append(layer)
else:
for child_layer in layer.children():
get_norm_layers(
child_layer, norm_layers_, share_layer_norm, share_group_norm
)
def register_norm_forward(
norm_layer: Union[nn.GroupNorm, nn.LayerNorm],
) -> Union[nn.GroupNorm, nn.LayerNorm]:
if not hasattr(norm_layer, "orig_forward"):
setattr(norm_layer, "orig_forward", norm_layer.forward)
orig_forward = norm_layer.orig_forward
def forward_(hidden_states: T) -> T:
n = hidden_states.shape[-2]
hidden_states = concat_first(hidden_states, dim=-2)
hidden_states = orig_forward(hidden_states) # type: ignore
return hidden_states[..., :n, :]
norm_layer.forward = forward_ # type: ignore
return norm_layer
def register_shared_norm(
model: ModelPatcher,
share_group_norm: bool = True,
share_layer_norm: bool = True,
):
norm_layers = {"group": [], "layer": []}
get_norm_layers(model.model, norm_layers, share_layer_norm, share_group_norm)
print(
f"Patching {len(norm_layers['group'])} group norms, {len(norm_layers['layer'])} layer norms."
)
return [register_norm_forward(layer) for layer in norm_layers["group"]] + [
register_norm_forward(layer) for layer in norm_layers["layer"]
]
SHARE_NORM_OPTIONS = ["both", "group", "layer", "disabled"]
SHARE_ATTN_OPTIONS = ["q+k", "q+k+v", "disabled"]
class StyleAlignedSampleReferenceLatents:
@classmethod
def INPUT_TYPES(s):
return {"required":
{
"reference_image": ("IMAGE",),
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING", ),
"model": ("MODEL",),
"vae": ("VAE", ),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"scheduler": (comfy.samplers.KSampler.SCHEDULERS.reverse(), ),
"denoise": ("FLOAT", {"default": 1, "min": 0.0, "max": 1.0, "step": 0.01}),
}
}
RETURN_TYPES = ("STEP_LATENTS","LATENT")
RETURN_NAMES = ("ref_latents", "noised_output")
FUNCTION = "run"
# CATEGORY = "style_aligned"
CATEGORY = "♾️Mixlab/Style"
def run(self, reference_image, positive, negative, model, vae, seed, steps, cfg,scheduler,denoise):
# TODO noise_mask?
def vae_encode_crop_pixels(pixels):
x = (pixels.shape[1] // 8) * 8
y = (pixels.shape[2] // 8) * 8
if pixels.shape[1] != x or pixels.shape[2] != y:
x_offset = (pixels.shape[1] % 8) // 2
y_offset = (pixels.shape[2] % 8) // 2
pixels = pixels[:, x_offset:x + x_offset, y_offset:y + y_offset, :]
return pixels
pixels=vae_encode_crop_pixels(reference_image)
t = vae.encode(pixels[:,:,:,:3])
latent_image = {"samples":t}
noise_seed=seed
sampler_name="ddim"
sampler = comfy.samplers.sampler_object(sampler_name)
total_steps = steps
if denoise < 1.0:
total_steps = int(steps/denoise)
comfy.model_management.load_models_gpu([model])
sigmas = comfy.samplers.calculate_sigmas_scheduler(model.model, scheduler, total_steps).cpu()
sigmas = sigmas[-(steps + 1):]
sigmas = sigmas.flip(0)
if sigmas[0] == 0:
sigmas[0] = 0.0001
latent = latent_image
latent_image = latent["samples"]
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
noise_mask = None
if "noise_mask" in latent:
noise_mask = latent["noise_mask"]
ref_latents = []
def callback(step: int, x0: T, x: T, steps: int):
ref_latents.insert(0, x[0])
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
samples = comfy.sample.sample_custom(model, noise, cfg, sampler, sigmas, positive, negative, latent_image, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=noise_seed)
out = latent.copy()
out["samples"] = samples
out_noised = out
ref_latents = torch.stack(ref_latents)
return (ref_latents, out_noised)
class StyleAlignedReferenceSampler:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"ref_latents": ("STEP_LATENTS",),
"reference_image_text": ("STRING", {"multiline": True}),
"model": ("MODEL",),
"clip": ("CLIP", ),
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING",),
"share_norm": (SHARE_NORM_OPTIONS,),
"share_attn": (SHARE_ATTN_OPTIONS,),
"scale": ("FLOAT", {"default": 1, "min": 0, "max": 2.0, "step": 0.01}),
"batch_size": ("INT", {"default": 2, "min": 1, "max": 8, "step": 1}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
},
}
RETURN_TYPES = ("LATENT", "LATENT")
RETURN_NAMES = ("output", "denoised_output")
FUNCTION = "patch"
# CATEGORY = "style_aligned"
CATEGORY = "♾️Mixlab/Style"
def patch(
self,
ref_latents,
reference_image_text,
model,
clip,
positive,
negative,
share_norm,
share_attn,
scale,
batch_size,
seed,steps,cfg,scheduler,denoise
) -> "tuple[dict, dict]":
m = model.clone()
# ref_latents = vae.encode(reference_image[:,:,:,:3])
tokens = clip.tokenize(reference_image_text)
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
ref_positive=[[cond, {"pooled_output": pooled}]]
noise_seed=seed
total_steps = steps
if denoise < 1.0:
total_steps = int(steps/denoise)
# comfy.model_management.load_models_gpu([model])
sigmas = comfy.samplers.calculate_sigmas_scheduler(model.model, scheduler, total_steps).cpu()
sigmas = sigmas[-(steps + 1):]
sampler_name="ddim"
sampler = comfy.samplers.sampler_object(sampler_name)
args = StyleAlignedArgs(share_attn)
# Concat batch with style latent
style_latent_tensor = ref_latents[0].unsqueeze(0)
height, width = style_latent_tensor.shape[-2:]
latent_t = torch.zeros(
[batch_size, 4, height, width], device=ref_latents.device
)
latent = {"samples": latent_t}
noise = comfy.sample.prepare_noise(latent_t, noise_seed)
latent_t = torch.cat((style_latent_tensor, latent_t), dim=0)
ref_noise = torch.zeros_like(noise[0]).unsqueeze(0)
noise = torch.cat((ref_noise, noise), dim=0)
x0_output = {}
preview_callback = latent_preview.prepare_callback(m, sigmas.shape[-1] - 1, x0_output)
# Replace first latent with the corresponding reference latent after each step
def callback(step: int, x0: T, x: T, steps: int):
preview_callback(step, x0, x, steps)
if (step + 1 < steps):
# 当ref_latents的step不够时
if step+1>len(ref_latents)-1:
step=len(ref_latents)-2
x[0] = ref_latents[step+1]
x0[0] = ref_latents[step+1]
# Register shared norms
share_group_norm = share_norm in ["group", "both"]
share_layer_norm = share_norm in ["layer", "both"]
register_shared_norm(m, share_group_norm, share_layer_norm)
# Patch cross attn
m.set_model_attn1_patch(SharedAttentionProcessor(args, scale))
# Add reference conditioning to batch
batched_condition = []
for i,condition in enumerate(positive):
additional = condition[1].copy()
batch_with_reference = torch.cat([ref_positive[i][0], condition[0].repeat([batch_size] + [1] * len(condition[0].shape[1:]))], dim=0)
if 'pooled_output' in additional and 'pooled_output' in ref_positive[i][1]:
# combine pooled output
pooled_output = torch.cat([ref_positive[i][1]['pooled_output'], additional['pooled_output'].repeat([batch_size]
+ [1] * len(additional['pooled_output'].shape[1:]))], dim=0)
additional['pooled_output'] = pooled_output
if 'control' in additional:
if 'control' in ref_positive[i][1]:
# combine control conditioning
control_hint = torch.cat([ref_positive[i][1]['control'].cond_hint_original, additional['control'].cond_hint_original.repeat([batch_size]
+ [1] * len(additional['control'].cond_hint_original.shape[1:]))], dim=0)
cloned_controlnet = additional['control'].copy()
cloned_controlnet.set_cond_hint(control_hint, strength=additional['control'].strength, timestep_percent_range=additional['control'].timestep_percent_range)
additional['control'] = cloned_controlnet
else:
# add zeros for first in batch
control_hint = torch.cat([torch.zeros_like(additional['control'].cond_hint_original), additional['control'].cond_hint_original.repeat([batch_size]
+ [1] * len(additional['control'].cond_hint_original.shape[1:]))], dim=0)
cloned_controlnet = additional['control'].copy()
cloned_controlnet.set_cond_hint(control_hint, strength=additional['control'].strength, timestep_percent_range=additional['control'].timestep_percent_range)
additional['control'] = cloned_controlnet
batched_condition.append([batch_with_reference, additional])
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
samples = comfy.sample.sample_custom(
m,
noise,
cfg,
sampler,
sigmas,
batched_condition,
negative,
latent_t,
callback=callback,
disable_pbar=disable_pbar,
seed=noise_seed,
)
# remove reference image
samples = samples[1:]
out = latent.copy()
out["samples"] = samples
if "x0" in x0_output:
out_denoised = latent.copy()
x0 = x0_output["x0"][1:]
out_denoised["samples"] = m.model.process_latent_out(x0.cpu())
else:
out_denoised = out
return (out, out_denoised)
class StyleAlignedBatchAlign:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"share_norm": (SHARE_NORM_OPTIONS,),
"share_attn": (SHARE_ATTN_OPTIONS,),
"scale": ("FLOAT", {"default": 1, "min": 0, "max": 1.0, "step": 0.1}),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "patch"
# CATEGORY = "style_aligned"
CATEGORY = "♾️Mixlab/Style"
def patch(
self,
model: ModelPatcher,
share_norm: str,
share_attn: str,
scale: float,
):
m = model.clone()
share_group_norm = share_norm in ["group", "both"]
share_layer_norm = share_norm in ["layer", "both"]
register_shared_norm(model, share_group_norm, share_layer_norm)
args = StyleAlignedArgs(share_attn)
m.set_model_attn1_patch(SharedAttentionProcessor(args, scale))
return (m,)
-437
View File
@@ -1,437 +0,0 @@
from transformers import pipeline, set_seed,AutoTokenizer, AutoModelForSeq2SeqLM
import random
import re
import os,sys
import folder_paths
# from PIL import Image
import importlib.util
import comfy.utils
# import numpy as np
import torch
import random
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")
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(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)
except ModuleNotFoundError:
return False
return spec is not None
try:
if is_installed('sentencepiece')==False:
import subprocess
# 安装
print('#pip install sentencepiece')
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'sentencepiece'], capture_output=True, text=True)
#检查命令执行结果
if result.returncode == 0 and is_installed('sentencepiece'):
print("#install success")
_available=True
else:
print("#install error")
_available=False
else:
_available=True
except:
_available=False
def translate(text):
global text_pipe,zh_en_model,zh_en_tokenizer
if zh_en_model==None:
zh_en_model = AutoModelForSeq2SeqLM.from_pretrained(zh_en_model_path).eval()
zh_en_tokenizer = AutoTokenizer.from_pretrained(zh_en_model_path,padding=True, truncation=True)
zh_en_model.to("cuda" if torch.cuda.is_available() else "cpu")
with torch.no_grad():
encoded = zh_en_tokenizer([text], return_tensors="pt")
encoded.to(zh_en_model.device)
sequences = zh_en_model.generate(**encoded)
return zh_en_tokenizer.batch_decode(sequences, skip_special_tokens=True)[0]
# input = "青春不能回头,所以青春没有终点。 ——《火影忍者》"
# print(input, translate(input))
def text_generate(text_pipe,input,seed=None):
if seed==None:
seed = random.randint(100, 1000000)
set_seed(seed)
for count in range(6):
sequences = text_pipe(input, max_length=random.randint(60, 90), num_return_sequences=8)
list = []
for sequence in sequences:
line = sequence['generated_text'].strip()
if line != input and len(line) > (len(input) + 4) and line.endswith((":", "-", "—")) is False:
list.append(line)
result = "\n".join(list)
result = re.sub('[^ ]+\.[^ ]+','', result)
result = result.replace("<", "").replace(">", "")
if result != "":
return result
if count == 5:
return result
# input = "Youth can't turn back, so there's no end to youth."
# print(input, text_generate(input))
import re
def correct_prompt_syntax(prompt=""):
# print("input prompt",prompt)
corrected_elements = []
# 处理成统一的英文标点
prompt = prompt.replace('(', '(').replace(')', ')').replace(',', ',').replace(';', ',').replace('。', '.').replace(':',':')
# 删除多余的空格
prompt = re.sub(r'\s+', ' ', prompt).strip()
prompt = prompt.replace("< ","<").replace(" >",">").replace("( ","(").replace(" )",")").replace("[ ","[").replace(' ]',']')
# 分词
prompt_elements = prompt.split(',')
def balance_brackets(element, open_bracket, close_bracket):
open_brackets_count = element.count(open_bracket)
close_brackets_count = element.count(close_bracket)
return element + close_bracket * (open_brackets_count - close_brackets_count)
for element in prompt_elements:
element = element.strip()
# 处理空元素
if not element:
continue
# 检查并处理圆括号、方括号、尖括号
if element[0] in '([':
corrected_element = balance_brackets(element, '(', ')') if element[0] == '(' else balance_brackets(element, '[', ']')
elif element[0] == '<':
corrected_element = balance_brackets(element, '<', '>')
else:
# 删除开头的右括号或右方括号
corrected_element = element.lstrip(')]')
corrected_elements.append(corrected_element)
# 重组修正后的prompt
return ','.join(corrected_elements)
# # 示例使用
# test_prompt = "((middle-century castles)), [forsaken: 0.8], (mystery dragons: 1.3, mist forests, sunsets, quiet; (((dummy)), [fisting city: 0.5] background, radiant, soft and flavoured,] promising mountains, ((starry: 1.6), [[crowds], [middle-century castle: urban landscapes of the future: 0.5], [yellow: bright sun: 0.7], overlooking"
# corrected_prompt = correct_prompt_syntax(test_prompt)
# print(corrected_prompt)
def detect_language(input_str):
# 统计中文和英文字符的数量
count_cn = count_en = 0
for char in input_str:
if '\u4e00' <= char <= '\u9fff':
count_cn += 1
elif char.isalpha():
count_en += 1
# 根据统计的字符数量判断主要语言
if count_cn > count_en:
return "cn"
elif count_en > count_cn:
return "en"
else:
return "unknow"
#定义Prompt文法
grammar = """
start: sentence
sentence: phrase ("," phrase)*
phrase: emphasis | weight | word | lora | embedding | schedule
emphasis: "(" sentence ")" -> emphasis
| "[" sentence "]" -> weak_emphasis
weight: "(" word ":" NUMBER ")"
schedule: "[" word ":" word ":" NUMBER "]"
lora: "<" WORD ":" WORD (":" NUMBER)? (":" NUMBER)? ">"
embedding: "embedding" ":" WORD (":" NUMBER)? (":" NUMBER)?
word: WORD
NUMBER: /\s*-?\d+(\.\d+)?\s*/
WORD: /[^,:\(\)\[\]<>]+/
"""
@v_args(inline=True) # Decorator to flatten the tree directly into the function arguments
class ChinesePromptTranslate(Transformer):
def sentence(self, *args):
return ", ".join(args)
def phrase(self, *args):
return "".join(args)
def emphasis(self, *args):
# Reconstruct the emphasis with translated content
return "(" + "".join(args) + ")"
def weak_emphasis(self, *args):
print('weak_emphasis:',args)
return "[" + "".join(args) + "]"
def embedding(self,*args):
print('prompt embedding',args[0])
if len(args) == 1:
# print('prompt embedding',str(args[0]))
# 只传递了一个参数,意味着只有embedding名称没有数字
embedding_name = str(args[0])
return f"embedding:{embedding_name}"
elif len(args) > 1:
embedding_name,*numbers = args
if len(numbers)==2:
return f"embedding:{embedding_name}:{numbers[0]}:{numbers[1]}"
elif len(numbers)==1:
return f"embedding:{embedding_name}:{numbers[0]}"
else:
return f"embedding:{embedding_name}"
def lora(self,*args):
print('lora prompt',*args)
if len(args) == 1:
return f"<lora:{loar_name}>"
elif len(args) > 1:
# print('lora', args)
_,loar_name,*numbers = args
loar_name = str(loar_name).strip()
if len(numbers)==2:
return f"<lora:{loar_name}:{numbers[0]}:{numbers[1]}>"
elif len(numbers)==1:
return f"<lora:{loar_name}:{numbers[0]}>"
else:
return f"<lora:{loar_name}>"
def weight(self, word,number):
translated_word = translate(str(word)).rstrip('.')
return f"({translated_word}:{str(number).strip()})"
def schedule(self,*args):
print('prompt schedule',args)
data = [str(arg).strip() for arg in args]
return f"[{':'.join(data)}]"
def word(self, word):
# Translate each word using the dictionary
if detect_language(str(word)) == "cn":
return translate(str(word)).rstrip('.')
else:
return str(word).rstrip('.')
class ChinesePrompt:
global _available
available=_available
@classmethod
def INPUT_TYPES(s):
return {"required": {
"text": ("STRING",{"multiline": True,"default": "", "dynamicPrompts": False}),
"generation": (["on","off"],{"default": "off"}),
},
"optional":{
"seed":("INT", {"default": 100, "min": 100, "max": 1000000}),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("prompt",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Prompt"
OUTPUT_NODE = True
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,)
global text_pipe,zh_en_model,zh_en_tokenizer
text_pipe= None
zh_en_model=None
zh_en_tokenizer=None
def run(self,text,seed,generation):
seed=seed[0]
generation=generation[0]
# 进度条
pbar = comfy.utils.ProgressBar(len(text)+1)
texts = [correct_prompt_syntax(t) for t in text]
global text_pipe,zh_en_model,zh_en_tokenizer
if zh_en_model==None:
zh_en_model = AutoModelForSeq2SeqLM.from_pretrained(zh_en_model_path).eval()
zh_en_tokenizer = AutoTokenizer.from_pretrained(zh_en_model_path,padding=True, truncation=True)
zh_en_model.to("cuda" if torch.cuda.is_available() else "cpu")
# zh_en_tokenizer.to("cuda" if torch.cuda.is_available() else "cpu")
text_pipe=pipeline('text-generation', model=text_generator_model_path,device="cuda" if torch.cuda.is_available() else "cpu")
# text_pipe.model.to("cuda" if torch.cuda.is_available() else "cpu")
prompt_result=[]
# print('zh_en_model device',zh_en_model.device,text_pipe.model.device,torch.cuda.current_device() )
en_texts=[]
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])
zh_en_model.to('cpu')
print("test en_text",en_texts)
# en_text.to("cuda" if torch.cuda.is_available() else "cpu")
pbar.update(1)
for t in en_texts:
if generation=='on':
prompt =text_generate(text_pipe,t,seed)
# 多条,还是单条
lines = prompt.split("\n")
longest_line = max(lines, key=len)
# print(longest_line)
prompt_result.append(longest_line)
else:
prompt_result.append(t)
pbar.update(1)
text_pipe.model.to('cpu')
print('prompt_result',prompt_result,)
# prompt_result = [','.join(correct_prompt_syntax(p)) for p in prompt_result]
if len(prompt_result)==0:
prompt_result=[""]
return {
"ui":{
"prompt": prompt_result
},
"result": (prompt_result,)}
class PromptGenerate:
global _available
available=_available
@classmethod
def INPUT_TYPES(s):
return {"required": {
"text": ("STRING",{"multiline": True,"default": "", "dynamicPrompts": False}),
},
"optional":{
"multiple": (["off","on"],),
"seed":("INT", {"default": 100, "min": 100, "max": 1000000}),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("prompt",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Prompt"
OUTPUT_NODE = True
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,)
global text_pipe
text_pipe= None
#
def run(self,text,multiple,seed):
global text_pipe
seed=seed[0]
multiple=multiple[0]
# 进度条
pbar = comfy.utils.ProgressBar(len(text))
text_pipe=pipeline('text-generation', model=text_generator_model_path,device="cuda" if torch.cuda.is_available() else "cpu")
prompt_result=[]
for t in text:
prompt =text_generate(text_pipe,t,seed)
prompt = prompt.split("\n")
if multiple=='off':
prompt = [max(prompt, key=len)]
for p in prompt:
prompt_result.append(p)
pbar.update(1)
text_pipe.model.to('cpu')
return {
"ui":{
"prompt": prompt_result
},
"result": (prompt_result,)}
-180
View File
@@ -1,180 +0,0 @@
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}}
+102 -560
View File
@@ -1,70 +1,10 @@
import os,platform
import re,random,json
import os
import re,random
from PIL import Image
import numpy as np
# FONT_PATH= os.path.abspath(os.path.join(os.path.dirname(__file__),'../assets/王汉宗颜楷体繁.ttf'))
import folder_paths
import matplotlib.font_manager as fm
import torch
import importlib.util
def create_incrementing_list(min_value, max_value, step, count):
l1 = [int(min_value + i * step) for i in range(count) if min_value + i * step <= max_value]
l2 = [float(min_value + i * step) for i in range(count) if min_value + i * step <= max_value]
return (l1,l2)
def split_list(lst, chunk_size, transition_size):
result = []
for i in range(0, len(lst), chunk_size):
start = i - transition_size
end = i + chunk_size + transition_size
result.append(lst[max(start, 0):end])
return result
def recursive_search(directory, excluded_dir_names=None):
if not os.path.isdir(directory):
return [], {}
if excluded_dir_names is None:
excluded_dir_names = []
result = []
dirs = {directory: os.path.getmtime(directory)}
for dirpath, subdirs, filenames in os.walk(directory, followlinks=True, topdown=True):
subdirs[:] = [d for d in subdirs if d not in excluded_dir_names]
for file_name in filenames:
relative_path = os.path.relpath(os.path.join(dirpath, file_name), directory)
result.append(relative_path)
for d in subdirs:
path = os.path.join(dirpath, d)
dirs[path] = os.path.getmtime(path)
return result, dirs
def filter_files_extensions(files, extensions):
return sorted(list(filter(lambda a: os.path.splitext(a)[-1].lower() in extensions or len(extensions) == 0, files)))
def get_system_font_path():
ps=[]
system = platform.system()
if system == "Windows":
ps.append(os.path.join(os.environ["WINDIR"], "Fonts"))
elif system == "Darwin":
ps.append(os.path.join("/Library", "Fonts"))
elif system == "Linux":
ps.append(os.path.join("/usr", "share", "fonts"))
ps.append(os.path.join("/usr", "local", "share", "fonts"))
ps=[p for p in ps if os.path.exists(p)]
file_paths=[]
for f in ps:
result, dirs=recursive_search(f)
for r in result:
file_paths.append(r)
file_paths=filter_files_extensions(file_paths,[".otf", ".ttf"])
return file_paths
# import json
# import hashlib
@@ -94,13 +34,13 @@ def create_temp_file(image):
) = folder_paths.get_save_image_path('tmp', output_dir)
im=tensor2pil(image)
image=tensor2pil(image)
image_file = f"{filename}_{counter:05}.png"
image_path=os.path.join(full_output_folder, image_file)
im.save(image_path,compress_level=4)
image.save(image_path,compress_level=4)
return [{
"filename": image_file,
@@ -120,65 +60,46 @@ def get_font_files(directory):
# 尝试获取系统字体
try:
font_paths = get_system_font_path()
for file in font_paths:
font_paths = fm.findSystemFonts()
for path in font_paths:
try:
font_name = os.path.splitext(file)[0]
font_path = file
font_files[font_name] = os.path.abspath(font_path)
font_prop = fm.FontProperties(fname=path)
font_name = font_prop.get_name()
font_files[font_name] = path
except Exception as e:
print(f"Error processing font {file}: {e}")
print(f"Error processing font {path}: {e}")
except Exception as e:
print(f"Error finding system fonts: {e}")
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)
def flatten_list(nested_list):
flat_list = []
for item in nested_list:
if isinstance(item, list):
flat_list.extend(flatten_list(item))
else:
if torch.is_tensor(item):
print('item.shape',item.shape)
for i in range(item.shape[0]):
flat_list.append(item[i:i + 1, ...])
else:
flat_list.append(item)
return flat_list
class ColorInput:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"color":("TCOLOR",),
},
}
RETURN_TYPES = ("STRING","INT","INT","INT","FLOAT",)
RETURN_NAMES = ("hex","r","g","b","a",)
RETURN_TYPES = ("STRING",)
# RETURN_NAMES = ("WIDTH","HEIGHT","X","Y",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Color"
CATEGORY = "♾️Mixlab/utils"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,False,False,False,False,)
OUTPUT_IS_LIST = (False,False,)
def run(self,color):
h=color['hex']
r=color['r']
g=color['g']
b=color['b']
a=color['a']
return (h,r,g,b,a,)
return (color,)
@@ -196,13 +117,13 @@ class FontInput:
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
CATEGORY = "♾️Mixlab/utils"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
OUTPUT_IS_LIST = (False,False,)
def run(self,font):
return (font_files[font],)
class TextToNumber:
@@ -211,17 +132,14 @@ class TextToNumber:
return {"required": {
"text": ("STRING",{"multiline": False,"default": "1"}),
"random_number": (["enable", "disable"],),
"max_num":("INT", {
"default": 10,
"min":2, #Minimum value
"number":("INT", {
"default": 0,
"min": 0, #Minimum value
"max": 10000000000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
},
"optional":{
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ("INT",)
@@ -229,12 +147,12 @@ class TextToNumber:
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Text"
CATEGORY = "♾️Mixlab/utils"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
def run(self,text,random_number,max_num,seed=0):
def run(self,text,random_number,number):
numbers = re.findall(r'\d+', text)
result=0
@@ -243,7 +161,7 @@ class TextToNumber:
# print(result)
if random_number=='enable' and result>0:
result= random.randint(1, max_num)
result= random.randint(1, 10000000000)
return {"ui": {"text": [text],"num":[result]}, "result": (result,)}
@@ -255,49 +173,26 @@ class FloatSlider:
"number":("FLOAT", {
"default": 0,
"min": 0, #Minimum value
"max": 0xffffffffffffffff, #Maximum value
"max": 1, #Maximum value
"step": 0.001, #Slider's step
"display": "slider" # Cosmetic only: display as "number" or "slider"
}),
"min_value":("FLOAT", {
"default": 0,
"min": -0xffffffffffffffff,
"max": 0xffffffffffffffff,
"step": 0.001,
"display": "number"
}),
"max_value":("FLOAT", {
"default": 1,
"min": -0xffffffffffffffff,
"max": 0xffffffffffffffff,
"step": 0.001,
"display": "number"
}),
"step":("FLOAT", {
"default": 0.001,
"min": -0xffffffffffffffff,
"max": 0xffffffffffffffff,
"step": 0.001,
"display": "number"
}),
},
}
RETURN_TYPES = ("FLOAT",)
RETURN_NAMES = ('FLOAT',)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
CATEGORY = "♾️Mixlab/utils"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
def run(self, number, min_value, max_value, step):
if number < min_value:
number = min_value
elif number > max_value:
number = max_value
return (number,)
def run(self,number):
return (number,)
class IntNumber:
@classmethod
@@ -307,30 +202,9 @@ class IntNumber:
"default": 0,
"min": -1, #Minimum value
"max": 0xffffffffffffffff,
"step": 1,
"display": "number"
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"min_value":("INT", {
"default": 0,
"min": -0xffffffffffffffff,
"max": 0xffffffffffffffff,
"step": 1,
"display": "number"
}),
"max_value":("INT", {
"default": 1,
"min": -0xffffffffffffffff,
"max": 0xffffffffffffffff,
"step": 1,
"display": "number"
}),
"step":("INT", {
"default": 1,
"min": -0xffffffffffffffff,
"max": 0xffffffffffffffff,
"step":1,
"display": "number"
}),
},
}
@@ -338,16 +212,13 @@ class IntNumber:
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
CATEGORY = "♾️Mixlab/utils"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
def run(self,number,min_value,max_value,step):
if number < min_value:
number= min_value
elif number > max_value:
number= max_value
def run(self,number):
return (number,)
class MultiplicationNode:
@@ -355,18 +226,11 @@ class MultiplicationNode:
def INPUT_TYPES(s):
return {"required": {
"numberA":(any_type,),
"multiply_by":("FLOAT", {
"default": 1,
"min": -2, #Minimum value
"max": 0xffffffffffffffff,
"step": 0.01, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"add_by":("FLOAT", {
"numberB":("FLOAT", {
"default": 0,
"min": -2000, #Minimum value
"min": -1, #Minimum value
"max": 0xffffffffffffffff,
"step": 0.01, #Slider's step
"step": 0.1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
})
},
@@ -376,21 +240,21 @@ class MultiplicationNode:
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Utils"
CATEGORY = "♾️Mixlab/utils"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,False,)
def run(self,numberA,multiply_by,add_by):
b=int(numberA*multiply_by+add_by)
a=float(numberA*multiply_by+add_by)
def run(self,numberA,numberB):
b=int(numberA*numberB)
a=float(numberA*numberB)
return (a,b,)
class TextInput:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"text": ("STRING",{"multiline": True,"default": ""})
"text": ("STRING",{"multiline": True,"default": ""}),
},
}
@@ -398,7 +262,7 @@ class TextInput:
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
CATEGORY = "♾️Mixlab/utils"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
@@ -407,61 +271,6 @@ class TextInput:
return (text,)
class IncrementingListNode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"min_value": ("FLOAT", {
"default": 0,
"min": -2000, #Minimum value
"max": 0xffffffffffffffff,
"step": 0.01, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"max_value": ("FLOAT", {
"default": 10,
"min": -2000, #Minimum value
"max": 0xffffffffffffffff,
"step": 0.01, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"step": ("FLOAT", {
"default": 0,
"min": -2000, #Minimum value
"max": 0xffffffffffffffff,
"step": 0.01, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"count": ("INT", {
"default": 1,
"min": 1, #Minimum value
"max": 0xffffffffffffffff,
"step":1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
})
},
"optional":{
"seed":("INT", {"default": -1, "min": -1, "max": 1000000}),
},
}
RETURN_TYPES = ("INT","FLOAT",)
RETURN_NAMES = ('int_list','float_list',)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Video"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (True,True,)
def run(self,min_value,max_value,step,count,seed):
print('create_incrementing_list',seed)
l1,l2=create_incrementing_list(min_value,max_value,step,count)
return (l1,l2,)
# 接收一个值,然后根据字符串或数值长度计算延迟时间,用户可以自定义延迟"字/s",延迟之后将转化
import comfy.samplers
@@ -494,7 +303,7 @@ class DynamicDelayProcessor:
},
"optional":{
"any_input":(any_type,),
"delay_by_text":("STRING",{"multiline":True,"dynamicPrompts": False,}),
"delay_by_text":("STRING",{"multiline":True,}),
"words_per_seconds":("FLOAT",{ "default":1.50,"min": 0.0,"max": 1000.00,"display":"Chars per second?"}),
"replace_output": (["disable","enable"],),
"replace_value":("INT",{ "default":-1,"min": 0,"max": 1000000,"display":"Replacement value"})
@@ -527,7 +336,7 @@ class DynamicDelayProcessor:
RETURN_TYPES = (any_type,)
RETURN_NAMES = ('output',)
CATEGORY = "♾️Mixlab/Utils"
CATEGORY = "♾️Mixlab/utils"
def run(self,any_input,delay_seconds,delay_by_text,words_per_seconds,replace_output,replace_value):
# print(f"Delay text:",delay_by_text )
# 获取开始时间戳
@@ -560,14 +369,14 @@ class AppInfo:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"name": ("STRING",{"multiline": False,"default": "Mixlab-App","dynamicPrompts": False}),
"input_ids":("STRING",{"multiline": True,"default": "\n".join(["1","2","3"]),"dynamicPrompts": False}),
"output_ids":("STRING",{"multiline": True,"default": "\n".join(["5","9"]),"dynamicPrompts": False}),
"name": ("STRING",{"multiline": False,"default": "Mixlab-App"}),
"image": ("IMAGE",),
"input_ids":("STRING",{"multiline": True,"default": "\n".join(["1","2","3"])}),
"output_ids":("STRING",{"multiline": True,"default": "\n".join(["5","9"])}),
},
"optional":{
"image": ("IMAGE",),
"description":("STRING",{"multiline": True,"default": "","dynamicPrompts": False}),
"description":("STRING",{"multiline": True,"default": ""}),
"version":("INT", {
"default": 1,
"min": 1,
@@ -575,60 +384,61 @@ class AppInfo:
"step": 1,
"display": "number"
}),
"share_prefix":("STRING",{"multiline": False,"default": "","dynamicPrompts": False}),
"link":("STRING",{"multiline": False,"default": "https://","dynamicPrompts": False}),
"category":("STRING",{"multiline": False,"default": "","dynamicPrompts": False}),
"auto_save": (["enable","disable"],),
}
}
RETURN_TYPES = ()
# RETURN_NAMES = ("IMAGE",)
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("IMAGE",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
OUTPUT_NODE = True
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):
name=name[0]
im=None
if image:
im=image[0][0]
#TODO batch 的方式需要处理
im=create_temp_file(im)
# image [img,] img[batch,w,h,a] 列表里面是batch,
def run(self,name,image,input_ids,output_ids,description,version):
input_ids=input_ids[0]
output_ids=output_ids[0]
description=description[0]
version=version[0]
share_prefix=share_prefix[0]
link=link[0]
category=category[0]
im=create_temp_file(image)
# id=get_json_hash([name,im,input_ids,output_ids,description,version])
return {"ui": {"json": [name,im,input_ids,output_ids,description,version,share_prefix,link,category]}, "result": ()}
return {"ui": {"json": [name,im,input_ids,output_ids,description,version]}, "result": (image,)}
class GetImageSize_:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
}
}
RETURN_TYPES = ("INT", "INT")
RETURN_NAMES = ("width", "height")
FUNCTION = "get_size"
CATEGORY = "♾️Mixlab/utils"
def get_size(self, image):
_, height, width, _ = image.shape
return (width, height)
class SwitchByIndex:
@classmethod
def INPUT_TYPES(cls):
return {
"optional":{
"A":(any_type,),
"B":(any_type,),
},
"required": {
"required": {
"A":(any_type,),
"B":(any_type,),
"index":("INT", {
"default": -1,
"min": -1,
@@ -636,76 +446,33 @@ class SwitchByIndex:
"step": 1,
"display": "number"
}),
"flat": (['off',"on"],),
}
}
RETURN_TYPES = (any_type,"INT",)
RETURN_NAMES = ("list", "count",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Utils"
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True, False,)
def run(self, A=[],B=[],index=-1,flat='on'):
flat=flat[0]
C=[]
index=index[0]
for a in A:
C.append(a)
for b in B:
C.append(b)
if flat=='on':
C=flatten_list(C)
if index>-1:
try:
C=[C[index]]
except Exception as e:
C=[C[-1]] #最后一个
return (C, len(C),)
class ListSplit:
@classmethod
def INPUT_TYPES(cls):
return {
"optional":{
"A":(any_type,),
},
"required": {
"chunk_size": ("INT", {"default": 10, "min": 1, "step": 1}),
"transition_size": ("INT", {"default": 0, "min": 0, "step": 1}),
"index": ("INT", {"default": -1, "min": -1, "step": 1}),
}
}
RETURN_TYPES = (any_type,)
RETURN_NAMES = ("B",)
RETURN_NAMES = ("C",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Utils"
CATEGORY = "♾️Mixlab/utils"
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,)
def run(self, A=[],chunk_size=[10],transition_size=[0],index=[-1]):
# print(len(A))
B=split_list(A,chunk_size[0],transition_size[0])
def run(self, A,B,index):
C=[]
index=index[0]
for a in A:
C.append(a)
for b in B:
C.append(b)
if index>-1:
try:
C=[C[index]]
except Exception as e:
C=[]
return (C,)
if index[0]>-1:
B=B[index[0]]
return (B,)
class LimitNumber:
@@ -736,7 +503,7 @@ class LimitNumber:
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
CATEGORY = "♾️Mixlab/utils"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
@@ -759,229 +526,4 @@ class LimitNumber:
return (nn,)
class ListStatistics:
@staticmethod
def count_types(lst):
type_count = {}
for item in lst:
item_type = type(item).__name__
if item_type not in type_count:
type_count[item_type] = []
if item_type in ['dict', 'str', 'int', 'float']:
type_count[item_type].append(item)
return type_count
# # 示例列表
# my_list = [1, 'hello', {'name': 'John'}, 3.14, {'age': 25}, 'world', 10]
# # 创建ListStatistics对象
# list_stats = ListStatistics()
# # 调用count_types方法进行统计
# result = list_stats.count_types(my_list)
# # 输出结果
# for item_type, values in result.items():
# print(item_type + ':')
# for value in values:
# print(value)
# print('---')
class TESTNODE_:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"ANY":(any_type,),
},
}
RETURN_TYPES = (any_type,)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Test"
OUTPUT_NODE = True
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,)
def run(self,ANY):
print(type(ANY))
try:
print(ANY[0].shape)
img= tensor2pil(ANY[0])
print(img.size)
except:
print('')
# data=ANY
list_stats = ListStatistics()
# 调用count_types方法进行统计
result = list_stats.count_types(ANY)
# 假设我们有一个模块文件名为 my_module.py,它位于 'importables' 目录下
module_path = os.path.join(os.path.dirname(__file__),'test.py')
# 使用 spec_from_file_location 获取模块的元数据(名称、定义等)
spec = importlib.util.spec_from_file_location('test', module_path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
functions = getattr(module, 'run') # 获取函数
functions(ANY)
return {"ui": {"data": result,"type":[str(type(ANY[0]))]}, "result": (ANY,)}
class TESTNODE_TOKEN:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"text":("STRING", {"forceInput": True,}),
"clip": ("CLIP", )
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Test"
OUTPUT_NODE = True
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
def run(self,text,clip=None):
# print(text)
tokens = clip.tokenize(text)
tokens=[v for v in tokens.values()][0][0]
tokens=json.dumps(tokens)
return (tokens,)
class CreateSeedNode:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ("INT",)
RETURN_NAMES = ("seed",)
OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Experiment"
def run(self, seed):
return (seed,)
class CreateCkptNames:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_names": ("STRING",{"multiline": True,"default": "\n".join(folder_paths.get_filename_list("checkpoints")),"dynamicPrompts": False}),
}
}
RETURN_TYPES = (any_type,)
RETURN_NAMES = ("ckpt_names",)
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (True,)
# OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Experiment"
def run(self, ckpt_names):
ckpt_names=ckpt_names.split('\n')
ckpt_names = [name for name in ckpt_names if name.strip()]
return (ckpt_names,)
class CreateLoraNames:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"lora_names": ("STRING",{"multiline": True,"default": "\n".join(folder_paths.get_filename_list("loras")),"dynamicPrompts": False}),
}
}
RETURN_TYPES = (any_type,"STRING",)
RETURN_NAMES = ("lora_names","prompt",)
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (True,True,)
# OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Experiment"
def run(self, lora_names):
lora_names=lora_names.split('\n')
lora_names = [name for name in lora_names if name.strip()]
prompts=[os.path.splitext(n)[0] for n in lora_names]
return (lora_names,prompts,)
class CreateSampler_names:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"sampler_names": ("STRING",{"multiline": True,"default": "\n".join(comfy.samplers.KSampler.SAMPLERS),"dynamicPrompts": False}),
}
}
RETURN_TYPES = (any_type,)
RETURN_NAMES = ("sampler_names",)
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (True,)
# OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Experiment"
def run(self, sampler_names):
sampler_names=sampler_names.split('\n')
sampler_names = [name for name in sampler_names if name.strip()]
return (sampler_names,)
+179
View File
@@ -0,0 +1,179 @@
# https://github.com/openai/consistencydecoder/blob/main/consistencydecoder/__init__.py
import folder_paths
from comfy import model_management
import math
import torch
import numpy as np
from PIL import Image
class ConsistencyDecoderWrapper:
def __init__(self, decoder):
self.decoder = decoder
def decode(self, x):
return self.decoder(x)
def _extract_into_tensor(arr, timesteps, broadcast_shape):
# from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L895 """
res = arr[timesteps].float()
dims_to_append = len(broadcast_shape) - len(res.shape)
return res[(...,) + (None,) * dims_to_append]
def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
# from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L45
betas = []
for i in range(num_diffusion_timesteps):
t1 = i / num_diffusion_timesteps
t2 = (i + 1) / num_diffusion_timesteps
betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta))
return torch.tensor(betas)
class ConsistencyDecoder:
def __init__(self, device="cuda:0", download_target=""):
self.n_distilled_steps = 64
# download_target = _download("https://openaipublic.azureedge.net/diff-vae/c9cebd3132dd9c42936d803e33424145a748843c8f716c0814838bdc8a2fe7cb/decoder.pt", download_root)
self.ckpt = torch.jit.load(download_target).to(device)
self.device = device
sigma_data = 0.5
betas = betas_for_alpha_bar(
1024, lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
).to(device)
alphas = 1.0 - betas
alphas_cumprod = torch.cumprod(alphas, dim=0)
self.sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / alphas_cumprod)
sigmas = torch.sqrt(1.0 / alphas_cumprod - 1)
self.c_skip = (
sqrt_recip_alphas_cumprod
* sigma_data**2
/ (sigmas**2 + sigma_data**2)
)
self.c_out = sigmas * sigma_data / (sigmas**2 + sigma_data**2) ** 0.5
self.c_in = sqrt_recip_alphas_cumprod / (sigmas**2 + sigma_data**2) ** 0.5
@staticmethod
def round_timesteps(
timesteps, total_timesteps, n_distilled_steps, truncate_start=True
):
with torch.no_grad():
space = torch.div(total_timesteps, n_distilled_steps, rounding_mode="floor")
rounded_timesteps = (
torch.div(timesteps, space, rounding_mode="floor") + 1
) * space
if truncate_start:
rounded_timesteps[rounded_timesteps == total_timesteps] -= space
else:
rounded_timesteps[rounded_timesteps == total_timesteps] -= space
rounded_timesteps[rounded_timesteps == 0] += space
return rounded_timesteps
@staticmethod
def ldm_transform_latent(z, extra_scale_factor=1):
channel_means = [0.38862467, 0.02253063, 0.07381133, -0.0171294]
channel_stds = [0.9654121, 1.0440036, 0.76147926, 0.77022034]
if len(z.shape) != 4:
raise ValueError()
z = z * 0.18215
channels = [z[:, i] for i in range(z.shape[1])]
channels = [
extra_scale_factor * (c - channel_means[i]) / channel_stds[i]
for i, c in enumerate(channels)
]
return torch.stack(channels, dim=1)
@torch.no_grad()
def __call__(
self,
features: torch.Tensor,
schedule=[1.0, 0.5],
):
features = self.ldm_transform_latent(features)
ts = self.round_timesteps(
torch.arange(0, 1024),
1024,
self.n_distilled_steps,
truncate_start=False,
)
shape = (
features.size(0),
3,
8 * features.size(2),
8 * features.size(3),
)
x_start = torch.zeros(shape, device=features.device, dtype=features.dtype)
schedule_timesteps = [int((1024 - 1) * s) for s in schedule]
for i in schedule_timesteps:
t = ts[i].item()
t_ = torch.tensor([t] * features.shape[0]).to(self.device)
noise = torch.randn_like(x_start)
x_start = (
_extract_into_tensor(self.sqrt_alphas_cumprod, t_, x_start.shape)
* x_start
+ _extract_into_tensor(
self.sqrt_one_minus_alphas_cumprod, t_, x_start.shape
)
* noise
)
c_in = _extract_into_tensor(self.c_in, t_, x_start.shape)
model_output = self.ckpt(c_in * x_start, t_, features=features)
B, C = x_start.shape[:2]
model_output, _ = torch.split(model_output, C, dim=1)
pred_xstart = (
_extract_into_tensor(self.c_out, t_, x_start.shape) * model_output
+ _extract_into_tensor(self.c_skip, t_, x_start.shape) * x_start
).clamp(-1, 1)
x_start = pred_xstart
return x_start
class VAELoader:
@classmethod
def INPUT_TYPES(s):
return {"required": { "vae_name": (folder_paths.get_filename_list("vae"), )}}
RETURN_TYPES = ("VAE",)
FUNCTION = "load_vae"
CATEGORY = "♾️Mixlab/_test"
#TODO: scale factor?
def load_vae(self, vae_name):
vae_path = folder_paths.get_full_path("vae", vae_name)
device = 'cuda:0'
# print('device',device)
consistencyDecoder = ConsistencyDecoder(device=device,
download_target=vae_path) # Model size: 2.49 GB
vae = ConsistencyDecoderWrapper(consistencyDecoder)
return (vae,)
class VAEDecode:
@classmethod
def INPUT_TYPES(s):
return {"required": { "samples": ("LATENT", ), "vae": ("VAE", )}}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "decode"
CATEGORY = "♾️Mixlab/_test"
def decode(self, vae, samples):
image = vae.decode(samples["samples"].to("cuda:0"))
image = image[0].cpu().numpy()
image = (image + 1.0) * 127.5
image = image.clip(0, 255).astype(np.uint8)
image = Image.fromarray(image.transpose(1, 2, 0))
image = image.convert("RGB")
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
return (image, )
-1001
View File
File diff suppressed because it is too large Load Diff
@@ -1,45 +0,0 @@
from comfy.ldm.modules.attention import default, optimized_attention, optimized_attention_masked
from .style_functions import adain, concat_first
class VisualStyleProcessor(object):
def __init__(self,
module_self,
keys_scale: float = 1.0,
enabled: bool = True,
adain_queries: bool = True,
adain_keys: bool = True,
adain_values: bool = False
):
self.module_self = module_self
self.keys_scale = keys_scale
self.enabled = enabled
self.adain_queries = adain_queries
self.adain_keys = adain_keys
self.adain_values = adain_values
def visual_style_forward(self, x, context, value, mask=None):
q = self.module_self.to_q(x)
context = default(context, x)
k = self.module_self.to_k(context)
if value is not None:
v = self.module_self.to_v(value)
del value
else:
v = self.module_self.to_v(context)
if self.enabled:
if self.adain_queries:
q = adain(q)
if self.adain_keys:
k = adain(k)
if self.adain_values:
v = adain(v)
k = concat_first(k, -2, self.keys_scale)
v = concat_first(v, -2)
if mask is None:
out = optimized_attention(q, k, v, self.module_self.heads)
else:
out = optimized_attention_masked(q, k, v, self.module_self.heads, mask)
return self.module_self.to_out(out)
@@ -1,60 +0,0 @@
import torch
from einops import rearrange
from dataclasses import dataclass
T = torch.Tensor
@dataclass(frozen=True)
class StyleAlignedArgs:
share_group_norm: bool = True
share_layer_norm: bool = True,
share_attention: bool = True
adain_queries: bool = True
adain_keys: bool = True
adain_values: bool = False
full_attention_share: bool = False
keys_scale: float = 1.
only_self_level: float = 0.
def expand_first(feat: T, scale=1., ) -> T:
b = feat.shape[0]
feat_style = torch.stack((feat[0], feat[b // 2])).unsqueeze(1)
if scale == 1:
feat_style = feat_style.expand(2, b // 2, *feat.shape[1:])
else:
feat_style = feat_style.repeat(1, b // 2, 1, 1, 1)
feat_style = torch.cat([feat_style[:, :1], scale * feat_style[:, 1:]], dim=1)
return feat_style.reshape(*feat.shape)
def concat_first(feat: T, dim=2, scale=1.) -> T:
feat_style = expand_first(feat, scale=scale)
return torch.cat((feat, feat_style), dim=dim)
def calc_mean_std(feat, eps: float = 1e-5) -> tuple[T, T]:
feat_std = (feat.var(dim=-2, keepdims=True) + eps).sqrt()
feat_mean = feat.mean(dim=-2, keepdims=True)
return feat_mean, feat_std
def adain(feat: T) -> T:
feat_mean, feat_std = calc_mean_std(feat)
feat_style_mean = expand_first(feat_mean)
feat_style_std = expand_first(feat_std)
feat = (feat - feat_mean) / feat_std
feat = feat * feat_style_std + feat_style_mean
return feat
def swapping_attention(key, value, chunk_size=2):
chunk_length = key.size()[0] // chunk_size # [text-condition, null-condition]
reference_image_index = [0] * chunk_length # [0 0 0 0 0]
key = rearrange(key, "(b f) d c -> b f d c", f=chunk_length)
key = key[:, reference_image_index] # ref to all
key = rearrange(key, "b f d c -> (b f) d c")
value = rearrange(value, "(b f) d c -> b f d c", f=chunk_length)
value = value[:, reference_image_index] # ref to all
value = rearrange(value, "b f d c -> (b f) d c")
return key, value
-172
View File
@@ -1,172 +0,0 @@
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)
-8
View File
@@ -1,8 +0,0 @@
import folder_paths
# 外挂一个文件,用来编写新的节点
def run(v):
output_dir = folder_paths.get_temp_directory()
print('1323',v,output_dir)
-38
View File
@@ -1,38 +0,0 @@
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
@@ -1,51 +0,0 @@
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
@@ -1,180 +0,0 @@
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
@@ -1,124 +0,0 @@
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
@@ -1,72 +0,0 @@
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
@@ -1,45 +0,0 @@
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
@@ -1,653 +0,0 @@
# 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
@@ -1,334 +0,0 @@
# 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)
@@ -1,219 +0,0 @@
# 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
@@ -1,218 +0,0 @@
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
@@ -1,475 +0,0 @@
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
@@ -1,10 +0,0 @@
{
"main_pass":
[
"-n", "-c:v", "libsvtav1",
"-pix_fmt", "yuv420p10le",
"-crf", "23"
],
"extension": "webm",
"environment": {"SVT_LOG": "1"}
}
-9
View File
@@ -1,9 +0,0 @@
{
"main_pass":
[
"-n", "-c:v", "libx264",
"-pix_fmt", "yuv420p",
"-crf", "19"
],
"extension": "mp4"
}
-11
View File
@@ -1,11 +0,0 @@
{
"main_pass":
[
"-n", "-c:v", "libx265",
"-pix_fmt", "yuv420p10le",
"-preset", "medium",
"-crf", "22",
"-x265-params", "log-level=quiet"
],
"extension": "mp4"
}
-9
View File
@@ -1,9 +0,0 @@
{
"main_pass":
[
"-n",
"-pix_fmt", "yuv420p",
"-crf", "23"
],
"extension": "webm"
}
-15
View File
@@ -1,15 +0,0 @@
[project]
name = "comfyui-mixlab-nodes"
description = "3D, ScreenShareNode & FloatingVideoNode, SpeechRecognition & SpeechSynthesis, GPT, LoadImagesFromLocal, Layers, Other Nodes, ..."
version = "0.32.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 = ""
+1 -14
View File
@@ -4,17 +4,4 @@ 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
torchaudio
soundfile>=0.12.1
# playwright
+215 -3021
View File
File diff suppressed because it is too large Load Diff
-3
View File
@@ -196,9 +196,6 @@ app.registerExtension({
}
if (bg) {
data.bg_image = await parseImage(bg)
if (!data.bg_image.match('data:image/')) {
delete data.bg_image
}
}
if (material) {
+50 -523
View File
@@ -2,48 +2,6 @@ import { app } from '../../../scripts/app.js'
import { $el } from '../../../scripts/ui.js'
import { api } from '../../../scripts/api.js'
import { td_bg } from './td_background.js'
console.log('td_bg', td_bg)
//本机安装的插件节点全集
window._nodesAll = null
//获取当前系统的插件,节点清单
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)
}
})
}
const base64Df =
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
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 get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 12 // the margin around the html element
@@ -70,21 +28,20 @@ function get_position_style (ctx, widget_width, y, node_height) {
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
flexDirection: 'row',
// alignItems: 'center',
justifyContent: 'flex-start',
zIndex: 9999999
justifyContent: 'flex-start'
}
}
async function drawImageToCanvas (imageUrl, sFactor = 320) {
async function drawImageToCanvas (imageUrl) {
var canvas = document.createElement('canvas')
var ctx = canvas.getContext('2d')
var img = new Image()
await new Promise((resolve, reject) => {
img.onload = function () {
var scaleFactor = sFactor / img.width
var scaleFactor = 320 / img.width
var canvasWidth = img.width * scaleFactor
var canvasHeight = img.height * scaleFactor
@@ -109,117 +66,26 @@ async function drawImageToCanvas (imageUrl, sFactor = 320) {
// 可以在这里执行其他操作,比如将Base64数据保存到服务器或显示在页面上
}
async function extractInputAndOutputData (
jsonData,
inputIds = [],
outputIds = []
) {
// workflow
// const workflow=jsonData.workflow;
// const nodes=workflow.nodes;
const data = jsonData.output
let input = []
let output = []
const seed = {}
const seedTitle = {}
function extractInputAndOutputData (jsonData, inputIds = [], outputIds = []) {
const data = jsonData
const input = []
const output = []
for (const id in data) {
if (data.hasOwnProperty(id)) {
let node = app.graph.getNodeById(id)
if (inputIds.includes(id)) {
// let node = app.graph.getNodeById(id)
let options = {}
let node = app.graph.getNodeById(id)
let options = []
// 模型
try {
if (node.type === 'CheckpointLoaderSimple') {
options = node.widgets.filter(w => w.name === 'ckpt_name')[0]
.options.values
} else if (node.type === 'LoraLoader') {
options = node.widgets.filter(w => w.name === 'lora_name')[0]
.options.values
}else if(node.type === 'LoraLoader'){
options =node.widgets.filter(w=>w.name==='lora_name')[0].options.values
}
} catch (error) {}
if (node.type == 'IntNumber' || node.type == 'FloatSlider') {
// min max step
let [v, min, max, step] = Array.from(node.widgets, w => w.value)
options = { min, max, step }
// node.widgets.filter(w => w.type === 'number')[0].options
}
if (node.type == 'PromptSlide') {
// min max step
options = node.widgets.filter(w => w.type === 'slider')[0].options
// 备选的keywords清单
try {
let keywords = node.widgets.filter(w => w.name === 'upload')[0]
.value
keywords = JSON.parse(keywords)
options.keywords = keywords
} catch (error) {
console.log(error)
}
}
if (node.type == 'ImagesPrompt_') {
//图库
// console.log('ImagesPrompt_', data[id])
let image_base64 = data[id].inputs.image_base64
let img_index = 0
let imgsData = JSON.parse(data[id].inputs.upload)
for (let index = 0; index < imgsData.length; index++) {
const imgd = imgsData[index].imgurl
imgsData[index].index = index
//TODO缩放大小
imgsData[index].imgurl = await parseImageToBase64(imgd)
if (image_base64 == imgsData[index].imgurl) {
img_index = index
}
}
options.images = imgsData
delete data[id].inputs.upload
delete data[id].inputs.image_base64
data[id].inputs.imageIndex = img_index
}
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]
if (output.links) {
// 有输出
options.hasMask = true
}
// loadImage的默认图,转为base64
let imgurl = app.graph.getNodeById(id).imgs[0].src + '&channel=rgb'
options.defaultImage = await drawImageToCanvas(imgurl, 512)
console.log('#loadImage的默认图', options)
}
input[inputIds.indexOf(id)] = {
...data[id],
title: node.title,
@@ -229,51 +95,14 @@ 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)
let node = app.graph.getNodeById(id)
// output.push()
output[outputIds.indexOf(id)] = {
...data[id],
title: node.title,
id,
options
}
}
if (
node.type === 'KSampler' ||
node.type == 'SamplerCustom' ||
node.type === 'ChinesePrompt_Mix' ||
node.type === 'Seed_'
) {
// seed 的类型收集
try {
seed[id] = node.widgets.filter(
w => w.name === 'seed' || w.name == 'noise_seed'
)[0].linkedWidgets[0].value
seedTitle[id] = node.title
} catch (error) {}
output[outputIds.indexOf(id)] = { ...data[id], title: node.title, id }
}
}
}
// 修复bug,当节点不存在时
input = input.filter(i => i)
output = output.filter(i => i)
return { input, output, seed, seedTitle }
return { input, output }
}
function getUrl () {
@@ -283,16 +112,6 @@ function getUrl () {
return url
}
const getLocalData = key => {
let data = {}
try {
data = JSON.parse(localStorage.getItem(key)) || {}
} catch (error) {
return {}
}
return data
}
async function save_app (json) {
let url = getUrl()
@@ -300,9 +119,7 @@ async function save_app (json) {
method: 'POST',
body: JSON.stringify({
data: json,
task: 'save_app',
filename: json.app.filename,
category: json.app.category
task: 'save_app'
})
})
return await res.json()
@@ -324,16 +141,9 @@ function downloadJsonFile (jsonData, fileName = 'mix_app.json') {
}, 0)
}
async function save (json, download = false, showInfo = true) {
let nodesAll = window._nodesAll || (await getObjectInfo())
console.log('####SAVE', nodesAll, json[0])
async function save (json, download = false) {
const name = json[0],
version = json[5],
share_prefix = json[6], //用于分享的功能扩展
link = json[7], //用于创建界面上的跳转链接
category = json[8] || '', //用于分类
description = json[4],
inputIds = json[2].split('\n').filter(f => f),
outputIds = json[3].split('\n').filter(f => f)
@@ -349,43 +159,18 @@ 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,
const { input, output } = extractInputAndOutputData(
data.output,
inputIds,
outputIds
)
let authorAvatar =
localStorage.getItem('_mixlab_author_avatar') || base64Df,
authorName =
localStorage.getItem('_mixlab_author_name') ||
localStorage.getItem('Comfy.userName'),
authorLink = localStorage.getItem('_mixlab_author_link') || ''
data.app = {
name,
description,
version,
input,
output,
seed, //控制是fixed 还是random
seedTitle,
share_prefix,
link,
category,
filename: `${name}_${version}.json`,
author: {
avatar: authorAvatar,
name: authorName,
link: authorLink
}
output
}
try {
@@ -393,91 +178,49 @@ async function save (json, download = false, showInfo = true) {
} catch (error) {}
// console.log(data.app)
// let http_workflow = app.graph.serialize()
await save_app(data)
if (download) {
await downloadJsonFile(data, data.app.filename)
}
if (showInfo) {
let open = window.confirm(
`You can now access the standalone application on a new page!\n${getUrl()}/mixlab/app?filename=${encodeURIComponent(
data.app.filename
)}&category=${encodeURIComponent(data.app.category)}`
if (download) {
await downloadJsonFile(
data,
`${data.app.name}_${data.app.version}_${new Date().toDateString()}.json`
)
if (open)
window.open(
`${getUrl()}/mixlab/app?filename=${encodeURIComponent(
data.app.filename
)}&category=${encodeURIComponent(data.app.category)}`
)
let open = window.confirm(
`You can now access the standalone application on a new page!\n${getUrl()}/mixlab/app?type=new`
)
if (open) window.open(`${getUrl()}/mixlab/app?type=new`)
} else {
await save_app(data)
let open = window.confirm(
`You can now access the standalone application on a new page!\n${getUrl()}/mixlab/app`
)
if (open) window.open(`${getUrl()}/mixlab/app`)
}
} catch (error) {
console.log('###error', error)
}
}
function getInputsAndOutputs () {
const inputs =
`LoadImage LoadImagesToBatch ImagesPrompt_ LoadAndCombinedAudio_ LoadVideoAndSegment_ VHS_LoadVideo CLIPTextEncode PromptSlide TextInput_ Color FloatSlider IntNumber CheckpointLoaderSimple LoraLoader`.split(
' '
),
outputs =
`SaveTripoSRMesh,PreviewImage,SaveImage,TransparentImage,ShowTextForGPT,CombineAudioVideo,VHS_VideoCombine,VideoCombine_Adv,Image Save,SaveImageAndMetadata_,ClipInterrogator`.split(
','
)
let inputsId = [],
outputsId = []
for (let node of app.graph._nodes) {
if (inputs.includes(node.type)) {
inputsId.push(node.id)
}
if (outputs.includes(node.type)) {
outputsId.push(node.id)
}
}
return {
input: inputsId,
output: outputsId
console.log('###SpeechRecognition', error)
}
}
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
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
// console.log('#orig_nodeCreated', this)
// 自动计算workflow里哪些节点支持
let input_ids = this.widgets.filter(w => w.name == 'input_ids')[0],
output_ids = this.widgets.filter(w => w.name == 'output_ids')[0]
const { input, output } = getInputsAndOutputs()
input_ids.value = input.join('\n')
output_ids.value = output.join('\n')
// console.log(this)
const widget = {
type: 'div',
name: 'AppInfoRun',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
{...get_position_style(
get_position_style(
ctx,
widget_width,
node.size[1] - widget_height,
node.widgets[4].last_y + 24,
node.size[1]
),zIndex:1}
)
)
}
}
@@ -493,7 +236,7 @@ app.registerExtension({
widget.div = $el('div', {})
const btn = document.createElement('button')
btn.innerText = 'Save & Open'
btn.innerText = 'Save For App'
btn.style = style
btn.addEventListener('click', () => {
@@ -503,7 +246,6 @@ app.registerExtension({
} else {
alert('Please run the workflow before saving')
// app.queuePrompt(0, 1)
this.widgets.filter(w => w.name === 'version')[0].value += 1
}
})
@@ -519,184 +261,13 @@ app.registerExtension({
} else {
alert('Please run the workflow before saving')
// app.queuePrompt(0, 1)
this.widgets.filter(w => w.name === 'version')[0].value += 1
}
})
//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`
let authorAvatar = document.createElement('img')
authorAvatar.className = `${'comfy-multiline-input'}`
authorAvatar.style = `outline: none;
border: none;
padding: 4px;
width: 32px;
cursor: pointer;
height: 32px;`
if (localStorage.getItem('_mixlab_author_avatar')) {
authorAvatar.src =
localStorage.getItem('_mixlab_author_avatar') || base64Df
}
let authorAvatarUpload = document.createElement('input')
authorAvatarUpload.type = 'file'
authorAvatarUpload.style = `display:none`
let authorAvatarInput = document.createElement('div')
authorAvatarInput.style = `display: flex;justify-content: flex-start;
align-items: center;`
let authorAvatarInputLabel = document.createElement('p')
authorAvatarInputLabel.innerText = 'Author Avatar'
authorAvatarInputLabel.className = `${'comfy-multiline-input'}`
authorAvatarInputLabel.style = `font-size:12px`
authorAvatar.addEventListener('click', e => {
authorAvatarUpload.click()
})
authorAvatarInputLabel.addEventListener('click', e => {
authorAvatarUpload.click()
})
authorAvatarUpload.addEventListener('change', event => {
const file = event.target.files[0]
const reader = new FileReader()
reader.onload = async e => {
let im = new Image()
im.src = e.target.result
authorAvatar.src = e.target.result
im.onload = () => {
let c = document.createElement('canvas')
let ctx = c.getContext('2d')
c.width = 72
c.height = 72
ctx.drawImage(
im,
0,
0,
im.naturalWidth,
im.naturalHeight,
0,
0,
c.width,
c.height
)
window._mixlab_author_avatar = c.toDataURL()
localStorage.setItem(
'_mixlab_author_avatar',
window._mixlab_author_avatar
)
}
}
// 以文本形式读取文件
reader.readAsDataURL(file)
})
author.appendChild(authorAvatarInput)
authorAvatarInput.appendChild(authorAvatarInputLabel)
authorAvatarInput.appendChild(authorAvatar)
authorAvatarInput.appendChild(authorAvatarUpload)
let authorName = document.createElement('input')
authorName.type = 'text'
authorName.value =
localStorage.getItem('_mixlab_author_name') ||
localStorage.getItem('Comfy.userName')
authorName.placeholder = 'author name'
authorName.className = `${'comfy-multiline-input'}`
authorName.style = `
outline: none;
border: none;
padding: 4px;
width: 100%;
cursor: pointer;
height: 32px;`
let authorNameInput = document.createElement('div')
authorNameInput.style = `display: flex;justify-content: flex-start;
align-items: center;`
let authorNameInputLabel = document.createElement('p')
authorNameInputLabel.innerText = 'Author Name'
authorNameInputLabel.className = `${'comfy-multiline-input'}`
authorNameInputLabel.style = `font-size:12px;width: 110px`
authorName.addEventListener('change', e => {
window._mixlab_author_name = authorName.value.trim()
localStorage.setItem(
'_mixlab_author_name',
window._mixlab_author_name
)
})
author.appendChild(authorNameInput)
authorNameInput.appendChild(authorNameInputLabel)
authorNameInput.appendChild(authorName)
// 社交链接
let authorLink = document.createElement('input')
authorLink.type = 'text'
authorLink.value = localStorage.getItem('_mixlab_author_link') || ''
authorLink.placeholder = 'author link'
authorLink.className = `${'comfy-multiline-input'}`
authorLink.style = `
outline: none;
border: none;
padding: 4px;
width: 100%;
cursor: pointer;
height: 32px;`
let authorLinkInput = document.createElement('div')
authorLinkInput.style = `display: flex;justify-content: flex-start;
align-items: center;`
let authorLinkInputLabel = document.createElement('p')
authorLinkInputLabel.innerText = 'Author Link'
authorLinkInputLabel.className = `${'comfy-multiline-input'}`
authorLinkInputLabel.style = `font-size:12px;width: 110px`
authorLink.addEventListener('change', e => {
window._mixlab_author_link = authorLink.value.trim()
localStorage.setItem(
'_mixlab_author_link',
window._mixlab_author_link
)
})
author.appendChild(authorLinkInput)
authorLinkInput.appendChild(authorLinkInputLabel)
authorLinkInput.appendChild(authorLink)
widget.div.appendChild(author)
let btns = document.createElement('div')
widget.div.appendChild(btns)
btns.appendChild(btn)
btns.appendChild(download)
btns.appendChild(tdBG)
document.body.appendChild(widget.div)
widget.div.appendChild(btn)
widget.div.appendChild(download)
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
@@ -706,66 +277,22 @@ app.registerExtension({
}
this.serialize_widgets = true //需要保存参数
window._mixlab_app_json = null
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
nodeType.prototype.onExecuted = async function (message) {
onExecuted?.apply(this, arguments)
console.log(message.json)
// console.log(this.widgets)
window._mixlab_app_json = message.json
try {
let a = this.widgets.filter(w => w.name === 'AppInfoRun')[0]
if (a) {
if (!a.value) a.value = 0
a.value += 1
}
const div = this.widgets.filter(w => w.div)[0].div
Array.from(div.querySelectorAll('button'), b =>
b.innerText != 'Canvas Mode' ? (b.style.background = 'yellow') : ''
Array.from(
div.querySelectorAll('button'),
b => (b.style.background = 'yellow')
)
} catch (error) {}
}
}
},
async loadedGraphNode (node, app) {
// console.log('#loadedGraphNode1111')
window._mixlab_app_json = null //切换workflow需要清空
if (node.type === 'AppInfo') {
let auto_save = node.widgets.filter(w => w.name == 'auto_save')[0]
if (auto_save) {
if (!['enable', 'disable'].includes(auto_save.value)) {
auto_save.value = 'enable'
}
}
// app.canvas.centerOnNode(node)
// app.canvas.setZoom(0.45)
}
}
})
api.addEventListener('execution_start', async ({ detail }) => {
console.log('#execution_start', detail)
window._mixlab_app_json = null
})
api.addEventListener('executed', async ({ detail }) => {
console.log('#executed', detail)
// window._mixlab_app_json=null;
const { output } = getInputsAndOutputs()
if (output.includes(parseInt(detail.node))) {
let appinfo = app.graph.findNodesByType('AppInfo')[0]
if (appinfo) {
let auto_save = appinfo.widgets.filter(w => w.name == 'auto_save')[0]
if (auto_save?.value === 'enable') {
// 自动保存
console.log('auto_save')
if (window._mixlab_app_json) save(window._mixlab_app_json, false, false)
}
}
}
})
-215
View File
@@ -396,218 +396,3 @@ 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)
}
}
}
})
-106
View File
@@ -1,106 +0,0 @@
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_ (url, messages, controller, callback) {
let request = await completion(url, messages, controller)
for await (const chunk of request) {
let content = chunk.data.choices[0].delta.content || ''
if (chunk.data.choices[0].role == 'assistant') {
//开始
content = ''
}
if (callback) callback(content)
}
}
+2 -8
View File
@@ -3,7 +3,7 @@ import { app } from '../../../scripts/app.js'
const repoOwner = 'shadowcz007' // 替换为仓库的所有者
const repoName = 'comfyui-mixlab-nodes' // 替换为仓库的名称
const version = 'v0.32.0'
const version = 'v0.6.0'
fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
.then(response => response.json())
@@ -17,13 +17,7 @@ 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(`<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>
app.ui.dialog.show(`<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;
-119
View File
@@ -1,119 +0,0 @@
import { app } from '../../../scripts/app.js'
// import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
function getRandomElements (arr, num) {
var result = []
var len = arr.length
for (var i = 0; i < num; i++) {
var randomIndex = Math.floor(Math.random() * len)
result.push(arr[randomIndex])
}
return result
}
const createPrompt = (node, prompts, items, sample) => {
const w = ComfyWidgets['STRING'](
node,
'text',
['STRING', { multiline: true }],
app
).widget
w.inputEl.readOnly = true
w.inputEl.style.opacity = 0.6
w.value = typeof prompts === 'string' ? prompts : prompts.join('\n\n')
const w2 = ComfyWidgets['STRING'](
node,
'text',
['STRING', { multiline: true }],
app
).widget
w2.inputEl.readOnly = true
w2.inputEl.style.opacity = 0.6
w2.value = typeof items === 'string' ? items : JSON.stringify(items, null, 2)
const w3 = ComfyWidgets['STRING'](
node,
'text',
['STRING', { multiline: true }],
app
).widget
w3.inputEl.readOnly = true
w3.inputEl.style.opacity = 0.6
w3.value = typeof sample === 'string' ? sample : sample.join('\n\n')
}
app.registerExtension({
name: 'Mixlab.prompt.ClipInterrogator',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeData.name === 'ClipInterrogator') {
function populate (prompts, items, random_samples) {
if (this.widgets) {
for (let i = 0; i < this.widgets.length; i++) {
if (this.widgets[i].type !== 'combo') this.widgets[i].onRemove?.()
}
this.widgets.length = 2
}
createPrompt(this, prompts, items, random_samples)
// console.log('ClipInterrogator', w, w2)
requestAnimationFrame(() => {
const sz = this.computeSize()
if (sz[0] < this.size[0]) {
sz[0] = this.size[0]
}
if (sz[1] < this.size[1]) {
sz[1] = this.size[1]
}
this.onResize?.(sz)
app.graph.setDirtyCanvas(true, false)
})
}
// When the node is executed we will be sent the input text, display this in the widget
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
// console.log('##', message)
populate.call(
this,
message.prompt,
message.analysis,
message.random_samples
)
}
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 === 'ClipInterrogator') {
try {
let widgets_values = node.widgets_values
console.log(widgets_values )
try {
if (widgets_values[2] && widgets_values[3] && widgets_values[4])
createPrompt(
node,
widgets_values[2],
widgets_values[3],
widgets_values[4]
)
} catch (error) {
console.log(error)
}
} catch (error) {}
}
}
})
+62 -207
View File
@@ -61,14 +61,14 @@ app.registerExtension({
async getCustomWidgets (app) {
return {
KEY (node, inputName, inputData, app) {
// console.log('##inputData', inputData)
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
return [128,32] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
let data = getLocalData('_mixlab_api_key')
@@ -199,224 +199,79 @@ app.registerExtension({
url[id] || 'https://api.openai.com/v1'
}
}
});
app.registerExtension({
name: 'Mixlab.GPT.SiliconflowLLM',
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.
},
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'SiliconflowLLM') {
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 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')
widget.div.appendChild(inputKey)
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
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 === 'SiliconflowLLM') {
let widget = node.widgets.filter(w => w.div)[0]
let apiKey = getLocalData('_mixlab_api_key');
let id = node.id
widget.div.querySelector('.Key').value = apiKey[id] || 'by Mixlab'
}
}
})
app.registerExtension({
name: 'Mixlab.GPT.ShowTextForGPT',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeData.name === 'ShowTextForGPT') {
function populate (text) {
text = text.filter(t => t && t?.trim())
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.name === "ShowTextForGPT") {
function populate(text) {
if (this.widgets) {
const pos = this.widgets.findIndex((w) => w.name === "text");
if (pos !== -1) {
for (let i = pos; i < this.widgets.length; i++) {
this.widgets[i].onRemove?.();
}
this.widgets.length = pos;
}
}
// console.log('ShowTextForGPT',text)
for (let list of text) {
const w = ComfyWidgets["STRING"](this, "text", ["STRING", { multiline: true }], app).widget;
w.inputEl.readOnly = true;
w.inputEl.style.opacity = 0.6;
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?.()
console.log('#ShowTextForGPT', this.widgets[i])
}
this.widgets.length = 2
}
for (let list of text) {
if (list) {
// console.log('#####', list)
const w = ComfyWidgets['STRING'](
this,
'show_text',
['STRING', { multiline: true }],
app
).widget
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)
data = Array.from(data, d => {
return {
...d,
content: decodeURIComponent(d.content)
}
})
list = JSON.stringify(data, null, 2)
try {
let data=JSON.parse(list);
data=Array.from(data,d=>{
return {
...d,
content:decodeURIComponent(d.content)
}
} catch (error) {
console.log(error)
}
w.value = list
})
list=JSON.stringify(data,null,2)
} catch (error) {
// console.log(error)
}
}
w.value =list;
}
// console.log('ShowTextForGPT',this.widgets.length)
requestAnimationFrame(() => {
if (this) {
const sz = this.computeSize()
if (sz[0] < this.size[0]) {
sz[0] = this.size[0]
}
if (sz[1] < this.size[1]) {
sz[1] = this.size[1]
}
this.onResize?.(sz)
app.graph.setDirtyCanvas(true, false)
}
})
}
requestAnimationFrame(() => {
const sz = this.computeSize();
if (sz[0] < this.size[0]) {
sz[0] = this.size[0];
}
if (sz[1] < this.size[1]) {
sz[1] = this.size[1];
}
this.onResize?.(sz);
app.graph.setDirtyCanvas(true, false);
});
}
// When the node is executed we will be sent the input text, display this in the widget
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
// console.log('##onExecuted', this, message)
if (message.text) populate.call(this, message.text)
}
// When the node is executed we will be sent the input text, display this in the widget
const onExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments);
populate.call(this, message.text);
};
const onConfigure = nodeType.prototype.onConfigure
nodeType.prototype.onConfigure = function () {
onConfigure?.apply(this, arguments)
if (this.widgets_values?.length) {
populate.call(this, this.widgets_values)
}
}
const onConfigure = nodeType.prototype.onConfigure;
nodeType.prototype.onConfigure = function () {
onConfigure?.apply(this, arguments);
if (this.widgets_values?.length) {
populate.call(this, this.widgets_values);
}
};
this.serialize_widgets = true //需要保存参数
}
},
async loadedGraphNode (node, app) {
if (node.type === 'ShowTextForGPT') {
let widget = node.widgets.filter(w => w.name == 'show_text')[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)
}
}
},
})
+5 -553
View File
@@ -1,39 +1,7 @@
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'
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();
@@ -60,9 +28,6 @@ async function uploadImage (blob, fileType = '.svg', filename) {
return src
}
const base64Df =
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
function base64ToBlobFromURL (base64URL, contentType) {
return fetch(base64URL).then(response => response.blob())
}
@@ -143,7 +108,7 @@ function createImage (url) {
})
}
const parseImageToBase64 = url => {
const parseImage = url => {
return new Promise((res, rej) => {
fetch(url)
.then(response => response.blob())
@@ -441,7 +406,9 @@ app.registerExtension({
this.serialize_widgets = true //需要保存参数
}
}
};
},
async loadedGraphNode (node, app) {
// Fires every time a node is constructed
@@ -475,518 +442,3 @@ app.registerExtension({
}
}
})
const createSelect = (imgDiv, select, opts, targetWidget, textWidget) => {
select.style.display = 'block'
let html = ''
let isMatch = false
for (const opt of opts) {
html += `<option value='${opt.keyword}' ${opt.selected ? 'selected' : ''}>${
opt.keyword
}</option>`
if (opt.selected) {
isMatch = true
imgDiv.src = opt.imgurl
// targetWidget.value = opt.keyword
}
}
select.innerHTML = html
if (!isMatch) {
// targetWidget.value = opts[0].keyword
imgDiv.src = opts[0].imgurl
}
// 添加change事件监听器
select.addEventListener('change', async function () {
// 获取选中的选项的值
var selectedOption = select.options[select.selectedIndex].value
let t = opts.filter(opt => opt.keyword === selectedOption)[0]
targetWidget.value = await parseImageToBase64(t.imgurl)
imgDiv.src = targetWidget.value
textWidget.value = t.keyword
})
// console.log(select)
}
app.registerExtension({
name: 'Mixlab.prompt.ImagesPrompt_',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'ImagesPrompt_') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = async function () {
orig_nodeCreated?.apply(this, arguments)
const image_prompt = this.widgets.filter(
w => w.name == 'image_base64'
)[0]
const image_text = this.widgets.filter(w => w.name == 'text')[0]
const node = this
const widget = {
type: 'div',
name: 'upload',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, y, node.size[1])
)
}
}
widget.div = $el('div', {})
// console.log('image_prompt',image_prompt)
const img = new Image()
img.src = image_prompt?.value || base64Df
widget.div.appendChild(img)
const btn = document.createElement('button')
btn.innerText = 'Upload Images JSON'
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;
`
const select = document.createElement('select')
select.style = `display:none;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: 100px;
`
widget.select = select
// const btn=document.createElement('button');
// btn.innerText='Upload'
btn.addEventListener('click', () => {
let inp = document.createElement('input')
inp.type = 'file'
inp.accept = '.json'
inp.click()
inp.addEventListener('change', event => {
// 获取选择的文件
// [{title,imageUrl}]
const file = event.target.files[0]
this.title = file.name.split('.')[0]
// console.log(file.name.split('.')[0])
// 创建文件读取器
const reader = new FileReader()
// 定义读取完成事件的回调函数
reader.onload = async event => {
// 读取完成后的文本内容
const json = JSON.parse(event.target.result)
console.log(node, json)
widget.value = JSON.stringify(json)
let img = widget.div.querySelector('img')
createSelect(img, select, json, image_prompt, image_text)
image_prompt.value = await parseImageToBase64(json[0].imgurl)
image_text.value = json[0].keyword
if (img) {
img.src = image_prompt.value
}
inp.remove()
}
// 以文本方式读取文件
reader.readAsText(file)
})
})
widget.div.appendChild(btn)
widget.div.appendChild(select)
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
if (this.onResize) {
this.onResize(this.size)
}
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'ImagesPrompt_') {
try {
let prompt = node.widgets.filter(w => w.name === 'image_base64')[0]
let text = node.widgets.filter(w => w.name === 'text')[0]
let uploadWidget = node.widgets.filter(w => w.name == 'upload')[0]
// console.log('##prompt',prompt.value)
let img = uploadWidget.div.querySelector('img')
let json = JSON.parse(uploadWidget.value)
for (let index = 0; index < json.length; index++) {
const j = json[index]
let base64 = await parseImageToBase64(j.imgurl)
if (base64 === prompt.value) {
json[index].selected = true
}
}
if (json && json[0]) {
uploadWidget.select.style.display = 'block'
createSelect(img, uploadWidget.select, json, prompt, text)
}
} catch (error) {}
}
}
})
const createInputImageForBatch = (base64, widget) => {
let im = new Image()
im.src = base64
im.style = `width: 88px;`
im.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
im.remove()
})
return im
}
app.registerExtension({
name: 'Mixlab.Comfy.LoadImagesToBatch',
async getCustomWidgets (app) {
return {
IMAGEBASE64 (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, 32] // 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 == 'LoadImagesToBatch') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
let imagesWidget = this.widgets.filter(w => w.name == 'images')[0]
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, 44, node.size[1])
)
},
serialize: false
}
widget.div = $el('div', {})
document.body.appendChild(widget.div)
let imagePreview = document.createElement('div')
let imagesDiv = document.createElement('div') //显示图片
imagesDiv.className = 'images_preview'
imagesDiv.style = `width: calc(100% - 14px);
display: flex;
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 => {
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)
}
reader.readAsDataURL(file)
})
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)
// document.addEventListener('wheel', handleMouseWheel)
const onRemoved = this.onRemoved
this.onRemoved = () => {
inputImage.remove()
widget.div.remove()
try {
// document.removeEventListener('wheel', handleMouseWheel)
} catch (error) {
console.log(error)
}
return onRemoved?.()
}
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]
let pre = imagePreview.div.querySelector('.images_preview')
for (const d of imagesWidget.value?.base64 || []) {
let im = createInputImageForBatch(d, imagesWidget)
pre.appendChild(im)
}
}
}
})
// 如何引入css
app.registerExtension({
name: 'Mixlab.output.ComparingTwoFrames_',
init () {
$el('link', {
rel: 'stylesheet',
href: '/extensions/comfyui-mixlab-nodes/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) {
Object.assign(
this.div.style,
get_position_style(ctx, 400, 44, node.size[1])
)
},
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
// }
// )
// }
}
}
})
+7 -591
View File
@@ -1,70 +1,8 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
// import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
function downloadJsonFile (jsonData, fileName = 'grid.json') {
const dataString = JSON.stringify(jsonData)
const blob = new Blob([dataString], { type: 'application/json' })
const url = URL.createObjectURL(blob)
const link = document.createElement('a')
link.href = url
link.download = fileName
link.click()
// 释放URL对象
setTimeout(() => {
URL.revokeObjectURL(url)
}, 0)
}
function createSelectWithOptions (options) {
const select = document.createElement('select')
options.forEach(option => {
const optionElement = document.createElement('option')
optionElement.text = option
optionElement.value = option
select.appendChild(optionElement)
})
select.style = `cursor: pointer;
font-weight: 300;
height: 30px;
min-width: 122px;
position: absolute;
top: 24px;
left: 88px;
z-index: 999999999999999;
`
return select
}
function drawCanvasWithText (w, h, tag, color = 'rgba(255,255,255,0.4)') {
const canvas = document.createElement('canvas')
const ctx = canvas.getContext('2d')
// 设置画布大小
canvas.width = w
canvas.height = h
// 绘制白色背景
ctx.fillStyle = color
ctx.fillRect(0, 0, canvas.width, canvas.height)
// 绘制文字
ctx.fillStyle = '#000000'
ctx.font = '20px Arial'
ctx.fillText(tag, 50, 50)
// 导出为Base64
const base64 = canvas.toDataURL()
return base64
}
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 4 // the margin around the html element
@@ -218,29 +156,6 @@ const parseSvg = async svgContent => {
return { data, image: base64, svgElement }
}
function findImages (nodeId) {
// 检查当前节点是否有 imgs 字段
const n = app.graph.getNodeById(nodeId)
if (n.imgs) {
return n.imgs
}
// 检查当前节点的 inputs 是否有 image 字段
if (n.inputs) {
for (let i = 0; i < n.inputs.length; i++) {
if (n.inputs[i].name === 'image' || n.inputs[i].name === 'images') {
// 获取新的 nodeId,并递归调用 findImages 函数
var linkId = n.inputs[i]?.link
var origin_id = app.graph.links[linkId].origin_id
return findImages(origin_id)
}
}
}
// 如果没有找到 imgs 字段或者 image 字段,则返回 null
return null
}
async function setArea (cw, ch, topBase64, base64, data, fn) {
let displayHeight = Math.round(window.screen.availHeight * 0.8)
let div = document.createElement('div')
@@ -412,196 +327,6 @@ async function setArea (cw, ch, topBase64, base64, data, fn) {
}
}
async function setAreaTags (cw, ch, grids, fn) {
let base64 = drawCanvasWithText(cw, ch, '', 'white')
let displayHeight = Math.round(window.screen.availHeight * 0.8)
let div = document.createElement('div')
div.innerHTML = `
<div id='ml_overlay' style='position: absolute;top:0;background: #251f1fc4;
height: 100vh;
z-index:999999;
width: 100%;'>
<img id='ml_video' style='position: absolute;
height: ${displayHeight}px;user-select: none;
-webkit-user-drag: none;
outline: 2px solid #eaeaea;
box-shadow: 8px 9px 17px #575757;' />
${Array.from(grids, g => {
const { label: tag, grid } = g
const [dx, dy, dw, dh] = grid
const base64Data = drawCanvasWithText(dw, dh, tag)
let x = 0,
y = 0,
width = (cw * displayHeight) / ch,
height = displayHeight
let imgWidth = cw
let imgHeight = ch
if (dw > 0 && dh > 0) {
// 相同尺寸窗口,恢复选区
x = (width * dx) / imgWidth
y = (height * dy) / imgHeight
width = (width * dw) / imgWidth
height = (height * dh) / imgHeight
}
return `<div class='ml_selection'
data-tag="${tag}"
style='position:absolute;
border: 2px dashed red;
pointer-events: none;
background-image: url("${base64Data}");
background-repeat: no-repeat;
background-size: cover;
left:${x}px;
top:${y}px;
width:${width}px;
height:${height}px;
'></div>`
})}
<div class="mx_close"> X </div>
</div>`
// document.body.querySelector('#ml_overlay')
document.body.appendChild(div)
const tags = Array.from(grids, g => g.label)
let select = createSelectWithOptions(tags)
document.body.appendChild(select)
let img = div.querySelector('#ml_video')
// let overlay = div.querySelector('#ml_overlay')
let selections = [...div.querySelectorAll('.ml_selection')]
let selection = selections.filter(
s => s.getAttribute('data-tag') === select.value
)[0]
select.addEventListener('change', e => {
selection = selections.filter(
s => s.getAttribute('data-tag') === select.value
)[0]
})
// console.log(select.value,selection)
let close = div.querySelector('.mx_close')
let startX, startY, endX, endY
let start = false
let setDone = false
// Set video source
img.src = base64
// canvas.toDataURL();
close.style = `cursor: pointer;
position: fixed;
left: 12px;
top: 12px;
z-index: 99999999;
background: black;
width: 44px;
height: 44px;
text-align: center;
line-height: 44px;`
// Add mouse events
img.addEventListener('mousedown', startSelection)
img.addEventListener('mousemove', updateSelection)
img.addEventListener('mouseup', endSelection)
const removeDiv = () => {
div.remove()
select?.remove()
close.removeEventListener('click', removeDiv)
img.removeEventListener('mousedown', startSelection)
img.removeEventListener('mousemove', updateSelection)
img.removeEventListener('mouseup', endSelection)
img.removeEventListener('mousedown', setDoneCheck)
}
close.addEventListener('click', removeDiv)
const setDoneCheck = event => {
console.log(setDone)
if (setDone) {
img.addEventListener('mousedown', startSelection)
img.addEventListener('mousemove', updateSelection)
img.addEventListener('mouseup', endSelection)
setDone = false
start = false
startX = event.clientX
startY = event.clientY
}
}
img.addEventListener('mousedown', setDoneCheck)
function remove () {
img.removeEventListener('mousedown', startSelection)
img.removeEventListener('mousemove', updateSelection)
img.removeEventListener('mouseup', endSelection)
setDone = true
// select?.remove()
}
function startSelection (event) {
if (start == false) {
startX = event.clientX
startY = event.clientY
updateSelection(event)
start = true
} else {
}
}
function updateSelection (event) {
endX = event.clientX
endY = event.clientY
// Calculate width, height, and coordinates
let width = Math.abs(endX - startX)
let height = Math.abs(endY - startY)
let left = Math.min(startX, endX)
let top = Math.min(startY, endY)
// Set selection style
selection.style.left = left + 'px'
selection.style.top = top + 'px'
selection.style.width = width + 'px'
selection.style.height = height + 'px'
}
function endSelection (event) {
endX = event.clientX
endY = event.clientY
// 获取img元素的真实宽度和高度
let imgWidth = img.naturalWidth
let imgHeight = img.naturalHeight
// 换算起始坐标
let realStartX = (startX / img.offsetWidth) * imgWidth
let realStartY = (startY / img.offsetHeight) * imgHeight
// 换算起始坐标
let realEndX = (endX / img.offsetWidth) * imgWidth
let realEndY = (endY / img.offsetHeight) * imgHeight
startX = realStartX
startY = realStartY
endX = realEndX
endY = realEndY
// Calculate width, height, and coordinates
let width = Math.round(Math.abs(endX - startX))
let height = Math.round(Math.abs(endY - startY))
let left = Math.round(Math.min(startX, endX))
let top = Math.round(Math.min(startY, endY))
if (width <= 0 && height <= 0) return remove()
if (!!fn) fn(select.value, left, top, width, height)
remove()
}
}
app.registerExtension({
name: 'Mixlab.layer.ShowLayer',
async getCustomWidgets (app) {
@@ -846,19 +571,15 @@ app.registerExtension({
}
}
try {
console.log('this.inputs', this.id)
let imgs = findImages(this.id)
// let topLinkId = this.inputs[0].link
// let topNodeId = app.graph.links[topLinkId].origin_id
let topIm = imgs[0]
console.log('this.inputs', this.inputs)
let topLinkId = this.inputs[0].link
let topNodeId = app.graph.links[topLinkId].origin_id
let topIm = app.graph.getNodeById(topNodeId).imgs[0]
let linkId = this.inputs[3].link
let nodeId = app.graph.links[linkId].origin_id
// console.log(linkId,this.inputs)
let imgs2 = findImages(nodeId)
let im = imgs2[0]
console.log(topIm, im)
let im = app.graph.getNodeById(nodeId).imgs[0]
// let src = im.src
setArea(
im.naturalWidth,
@@ -868,9 +589,7 @@ app.registerExtension({
data,
updateValue
)
} catch (error) {
console.log(error)
}
} catch (error) {}
})
}
}
@@ -890,306 +609,3 @@ app.registerExtension({
}
}
})
app.registerExtension({
name: 'Mixlab.layer.GridInput',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'GridInput') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = async function () {
orig_nodeCreated?.apply(this, arguments)
const grids_widget = this.widgets.filter(w => w.name == 'grids')[0]
const widget = {
type: 'div',
name: 'upload',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, y, node.size[1]),
{
justifyContent: 'flex-start'
}
)
}
}
widget.div = $el('div', {})
const addBtn = document.createElement('button')
addBtn.innerText = 'Add Box'
addBtn.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;
`
const vbtn = document.createElement('button')
vbtn.innerText = 'Set Box'
vbtn.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;
`
const btn = document.createElement('button')
btn.innerText = 'Upload JSON'
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;
`
addBtn.addEventListener('click', () => {
const { width, height, grids } = JSON.parse(grids_widget.value)
grids.push({
label: 'background',
grid: [12, 12, width - 24, height - 24]
})
grids_widget.value = JSON.stringify(
{
width,
height,
grids
},
null,
2
)
})
vbtn.addEventListener('click', () => {
const { width, height, grids } = JSON.parse(grids_widget.value)
setAreaTags(width, height, grids, (tag, x, y, w, h) => {
grids_widget.value = JSON.stringify(
{
width,
height,
grids: Array.from(grids, g => {
if (g.label === tag) {
g.grid = [x, y, w, h]
}
return g
})
},
null,
2
)
})
})
btn.addEventListener('click', () => {
let inp = document.createElement('input')
inp.type = 'file'
inp.accept = '.json'
inp.click()
inp.addEventListener('change', event => {
// 获取选择的文件
const file = event.target.files[0]
this.title = file.name.split('.')[0]
// console.log(file.name.split('.')[0])
// 创建文件读取器
const reader = new FileReader()
// 定义读取完成事件的回调函数
reader.onload = event => {
// 读取完成后的文本内容
const fileContent = JSON.parse(event.target.result)
const grids = fileContent
grids_widget.value = JSON.stringify(grids, null, 2)
// widget.value = grids
inp.remove()
}
// 以文本方式读取文件
reader.readAsText(file)
})
})
widget.div.appendChild(addBtn)
widget.div.appendChild(vbtn)
widget.div.appendChild(btn)
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
const r = onExecuted?.apply?.(this, arguments)
let json = message.json
if (json) {
json = {
width: json[0],
height: json[1],
grids: json[2]
}
grids_widget.value = JSON.stringify(json, null, 2)
// widget.value = json
}
return r
}
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
if (this.onResize) {
this.onResize(this.size)
}
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'GridInput') {
try {
const grids_widget = node.widgets.filter(w => w.name == 'grids')[0]
const { width, height, grids } = JSON.parse(grids_widget.value)
console.log('#GridInput', node, grids)
const div = node.widgets.filter(w => w.name == 'upload')[0]
div.div.querySelector('select').innerHTML = Array.from(
grids,
g => `<option value="${g.label}">${g.label}</option>`
).join('')
} catch (error) {}
}
}
})
app.registerExtension({
name: 'Mixlab.layer.GridDisplayAndSave',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'GridDisplayAndSave') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = async function () {
orig_nodeCreated?.apply(this, arguments)
const grids_widget = this.widgets.filter(w => w.name == 'grids')[0]
console.log('GridDisplayAndSave', grids_widget)
const widget = {
type: 'div',
name: 'save_json',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, y, node.size[1]),
{
justifyContent: 'flex-start',
flexDirection: 'column'
}
)
}
}
widget.div = $el('div', {})
const btn = document.createElement('button')
btn.innerText = 'Save JSON'
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;
max-width: 122px;
`
btn.addEventListener('click', () => {
if (window._mixlab_grid)
downloadJsonFile(
window._mixlab_grid,
this.widgets.filter(w => w.name == 'filename_prefix')[0]?.value +
'_grid.json'
)
})
widget.div.appendChild(btn)
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
const r = onExecuted?.apply?.(this, arguments)
let save_json = this.widgets.filter(d => d.name == 'save_json')[0]
let div = save_json?.div
// console.log('Test',message)
let image = message.image[0]
let json = message.json
if (image) {
const { filename, subfolder, type } = image
if (!div.querySelector('img')) {
let im = new Image()
div.appendChild(im)
im.style.width = '100%'
}
div.querySelector('img').src = api.apiURL(
`/view?filename=${encodeURIComponent(
filename
)}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
)
window._mixlab_grid = {
width: json[0],
height: json[1],
grids: json[2]
}
// console.log(src)
}
this.onResize?.(this.size)
return r
}
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
if (this.onResize) {
this.onResize(this.size)
}
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'GridDisplayAndSave') {
try {
let grids_widget = node.widgets.filter(w => w.name === 'grids')[0]
// let ks = getLocalData(`_mixlab_PromptSlide`)
let uploadWidget = node.widgets.filter(w => w.name == 'upload')[0]
// console.log('##widget', uploadWidget.value)
let grids = JSON.parse(uploadWidget.value)
} catch (error) {}
}
}
})
+16 -31
View File
@@ -1267,7 +1267,7 @@ app.registerExtension({
})
widget.PictureInPicture = $el('button', {
innerText: 'Picture In Picture',
innerText: 'PictureInPicture',
style: {
display: 'pictureInPictureEnabled' in document ? 'block' : 'none',
cursor: 'pointer',
@@ -1331,13 +1331,7 @@ app.registerExtension({
let w = 360,
s = widget.preview.videoWidth / widget.preview.videoHeight,
h = w / s || w
// console.log(h)
if (!window.documentPictureInPicture) {
window.alert(
'This feature is available only in secure contexts (HTTPS), in some or all supporting browsers. https://developer.mozilla.org/en-US/docs/Web/API/Document_Picture-in-Picture_API'
)
}
console.log(h)
const pipWindow = await documentPictureInPicture.requestWindow({
width: w,
@@ -1806,19 +1800,15 @@ const updateUI = node => {
pw.inputEl.title = `Total of ${prompts.length} prompts`
} else {
// 动态添加
// console.log('ComfyWidgets',ComfyWidgets.STRING(
// node,
// 'prompts',
// ['STRING', { multiline: true }]
// ))
// ComfyWidgets.STRING(this, "", ["", {default:this.properties.text, multiline: true}], app)
console.log('ComfyWidgets',ComfyWidgets.STRING(
node,
'prompts',
['STRING', { multiline: true }]
))
const w = ComfyWidgets.STRING(
node,
'prompts',
['STRING', { multiline: true }],
app
['STRING', { multiline: true }]
).widget
w.inputEl.readOnly = true
w.inputEl.style.opacity = 0.6
@@ -2099,13 +2089,13 @@ const node = {
name: 'RandomPrompt',
async init (app) {
// Any initial setup to run as soon as the page loads
// console.log('[logging]', 'extension init')
console.log('[logging]', 'extension init')
if (window.location.href.match('/?')) {
const { workflow } = getURLParameters(window.location.href)
if (workflow)
get_my_workflow().then(data => {
// console.log('#get_my_workflow', data)
console.log('#get_my_workflow', data)
let my_workflow = data.filter(
d => d.filename == 'my_workflow.json'
)[0]
@@ -2141,15 +2131,10 @@ const node = {
// }
},
loadedGraphNode (node, app) {
if (node.type === 'RandomPrompt') {
try {
let max_count = node.widgets.filter(w => w.name === 'max_count')[0]
max_count.value = node.widgets_values[0]
// console.log('RandomPrompt',max_count,node.widgets_values[0])
} catch (error) {
console.log(error)
}
}
// Fires for each node when loading/dragging/etc a workflow json or png
// If you break something in the backend and want to patch workflows in the frontend
// This is the place to do this
// console.log("[logging]", "loaded graph node: ", exportGraph(node.graph));
},
async nodeCreated (node) {
if (node.type === 'RandomPrompt') {
@@ -2242,7 +2227,7 @@ const node = {
const r = onExecuted?.apply?.(this, arguments)
let prompts = message.prompts
// console.log('executed', message)
console.log('executed', message)
// console.log('#RandomPrompt', this.widgets)
const pw = this.widgets.filter(w => w.name === 'prompts')[0]
@@ -2253,7 +2238,7 @@ const node = {
} else {
// 动态添加
const w = ComfyWidgets.STRING(
this,
node,
'prompts',
['STRING', { multiline: true }],
app
-565
View File
@@ -1,565 +0,0 @@
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 PhotoSwipeLightbox from '/extensions/comfyui-mixlab-nodes/lib/photoswipe-lightbox.esm.min.js'
function loadCSS (url) {
var link = document.createElement('link')
link.rel = 'stylesheet'
link.type = 'text/css'
link.href = url
document.getElementsByTagName('head')[0].appendChild(link)
// Create a style element
const style = document.createElement('style')
// Define the CSS rule for scrollbar width
const cssRule = `.pswp__custom-caption {
background: rgb(20 27 70);
font-size: 16px;
color: #fff;
width: calc(100% - 32px);
max-width: 980px;
padding: 2px 8px;
border-radius: 4px;
position: absolute;
left: 50%;
bottom: 16px;
transform: translateX(-50%);
}
.pswp__custom-caption a {
color: #fff;
text-decoration: underline;
}
.hidden-caption-content {
display: none;
}`
// Add the CSS rule to the style element
style.appendChild(document.createTextNode(cssRule))
// Append the style element to the document head
document.head.appendChild(style)
}
loadCSS('/extensions/comfyui-mixlab-nodes/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')
})
lightbox.on('uiRegister', function () {
lightbox.pswp.ui.registerElement({
name: 'custom-caption',
order: 9,
isButton: false,
appendTo: 'root',
html: 'Caption text',
onInit: (el, pswp) => {
lightbox.pswp.on('change', () => {
const currSlideElement = lightbox.pswp.currSlide.data.element
let captionHTML = ''
if (currSlideElement) {
const hiddenCaption = currSlideElement.querySelector(
'.hidden-caption-content'
)
if (hiddenCaption) {
// get caption from element with class hidden-caption-content
captionHTML = hiddenCaption.innerHTML
} else {
// get caption from alt attribute
captionHTML = currSlideElement
.querySelector('img')
.getAttribute('alt')
}
}
el.innerHTML = captionHTML || ''
})
}
})
})
lightbox.init()
}
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 - 24}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
paddingLeft: '12px',
display: 'flex',
flexDirection: 'row',
// alignItems: 'center',
justifyContent: 'space-between'
}
}
function createImage (url) {
let im = new Image()
return new Promise((res, rej) => {
im.onload = () => res(im)
im.src = url
})
}
async function fetchImage (url) {
try {
const response = await fetch(url)
const blob = await response.blob()
return blob
} catch (error) {
console.error('出现错误:', error)
}
}
const getLocalData = key => {
let data = {}
try {
data = JSON.parse(localStorage.getItem(key)) || {}
} catch (error) {
return {}
}
return data
}
const setLocalDataOfWin = (key, value) => {
localStorage.setItem(key, JSON.stringify(value))
// window[key] = value
}
const createSelect = (select, opts, targetWidget) => {
select.style.display = 'block'
let html = ''
let isMatch = false
for (const opt of opts) {
html += `<option value='${opt}' ${
targetWidget.value === opt ? 'selected' : ''
}>${opt}</option>`
if (targetWidget.value === opt) isMatch = true
}
select.innerHTML = html
if (!isMatch) targetWidget.value = opts[0]
// 添加change事件监听器
select.addEventListener('change', function () {
// 获取选中的选项的值
var selectedOption = select.options[select.selectedIndex].value
targetWidget.value = selectedOption
// console.log(widget,selectedOption)
})
}
app.registerExtension({
name: 'Mixlab.prompt.RandomPrompt',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'RandomPrompt') {
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]
// console.log('PromptSlide nodeData', prompt_keyword)
const widget = {
type: 'div',
name: 'upload',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, y, node.size[1])
)
}
}
widget.div = $el('div', {})
const btn = document.createElement('button')
btn.innerText = 'Upload Keywords'
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;
`
// const btn=document.createElement('button');
// btn.innerText='Upload'
btn.addEventListener('click', () => {
let inp = document.createElement('input')
inp.type = 'file'
inp.accept = '.txt'
inp.click()
inp.addEventListener('change', event => {
// 获取选择的文件
const file = event.target.files[0]
this.title = file.name.split('.')[0]
// console.log(file.name.split('.')[0])
// 创建文件读取器
const reader = new FileReader()
// 定义读取完成事件的回调函数
reader.onload = event => {
// 读取完成后的文本内容
const fileContent = event.target.result.split('\n')
const keywords = Array.from(fileContent, f => f.trim()).filter(
f => f
)
// 打印文件内容
// console.log(keywords)
mutable_prompt.value = keywords.join('\n')
inp.remove()
}
// 以文本方式读取文件
reader.readAsText(file)
})
})
widget.div.appendChild(btn)
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
if (this.onResize) {
this.onResize(this.size)
}
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'RandomPrompt') {
}
}
})
app.registerExtension({
name: 'Mixlab.prompt.PromptSlide',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'PromptSlide') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = async function () {
orig_nodeCreated?.apply(this, arguments)
const prompt_keyword = this.widgets.filter(
w => w.name == 'prompt_keyword'
)[0]
// console.log('PromptSlide nodeData', prompt_keyword)
const widget = {
type: 'div',
name: 'upload',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, y, node.size[1])
)
}
}
widget.div = $el('div', {})
const btn = document.createElement('button')
btn.innerText = 'Upload Keywords'
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;
`
const select = document.createElement('select')
select.style = `display:none;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: 100px;
`
widget.select = select
// const btn=document.createElement('button');
// btn.innerText='Upload'
btn.addEventListener('click', () => {
let inp = document.createElement('input')
inp.type = 'file'
inp.accept = '.txt'
inp.click()
inp.addEventListener('change', event => {
// 获取选择的文件
const file = event.target.files[0]
this.title = file.name.split('.')[0]
// console.log(file.name.split('.')[0])
// 创建文件读取器
const reader = new FileReader()
// 定义读取完成事件的回调函数
reader.onload = event => {
// 读取完成后的文本内容
const fileContent = event.target.result.split('\n')
const keywords = Array.from(fileContent, f => f.trim()).filter(
f => f
)
// 打印文件内容
// console.log(keywords)
widget.value = JSON.stringify(keywords)
// let ks = getLocalData(`_mixlab_PromptSlide`)
// ks[this.id] = keywords
// setLocalDataOfWin(`_mixlab_PromptSlide`, ks)
createSelect(select, keywords, prompt_keyword)
inp.remove()
}
// 以文本方式读取文件
reader.readAsText(file)
})
})
widget.div.appendChild(btn)
widget.div.appendChild(select)
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
if (this.onResize) {
this.onResize(this.size)
}
this.serialize_widgets = true //需要保存参数
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'PromptSlide') {
try {
let prompt = node.widgets.filter(w => w.name === 'prompt_keyword')[0]
// let ks = getLocalData(`_mixlab_PromptSlide`)
let uploadWidget = node.widgets.filter(w => w.name == 'upload')[0]
// console.log('##widget', uploadWidget.value)
let keywords = JSON.parse(uploadWidget.value)
// console.log('keywords',keywords)
let widget = node.widgets.filter(w => w.select)[0]
if (keywords && keywords[0]) {
widget.select.style.display = 'block'
createSelect(widget.select, keywords, prompt)
}
} catch (error) {}
}
}
})
const _createResult = async (node, widget, message) => {
widget.div.innerHTML = ``
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]
for (const img of imgs) {
let url = api.apiURL(
`/view?filename=${encodeURIComponent(img.filename)}&type=${
img.type
}&subfolder=${
img.subfolder
}${app.getPreviewFormatParam()}${app.getRandParam()}`
)
let image = await createImage(url)
// 创建card
let div = document.createElement('div')
div.className = 'card'
div.draggable = true
div.ondragend = async event => {
console.log('拖动停止')
let url = div.querySelector('img').src
let blob = await fetchImage(url)
let imageNode = null
// No image node selected: add a new one
if (!imageNode) {
const newNode = LiteGraph.createNode('LoadImage')
newNode.pos = [...app.canvas.graph_mouse]
imageNode = app.graph.add(newNode)
app.graph.change()
}
// const blob = item.getAsFile();
imageNode.pasteFile(blob)
}
div.setAttribute('data-scale', image.naturalHeight / image.naturalWidth)
let h = (image.naturalHeight * width) / image.naturalWidth
if (index % 2 === 0) height_add += h
div.style = `width: ${width}px;height:${h}px;position: relative;margin: 4px;`
div.innerHTML = `<a href="${url}"
data-pswp-width="${image.naturalWidth}"
data-pswp-height="${image.naturalHeight}"
target="_blank">
<img src="${url}" style='width: 100%' alt="${message.prompts[index]}"/>
</a>
<p style="position: absolute;
bottom: 0;
left: 0;
opacity: 0.6;
background-color: var(--comfy-input-bg);
color: var(--descrip-text);
margin: 0;
font-size: 12px;
padding: 5px;
text-align: left;">${message.prompts[index]}</p>`
widget.div.appendChild(div)
}
}
node.size[1] = 98 + height_add
}
app.registerExtension({
name: 'Mixlab.prompt.PromptImage',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'PromptImage') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
console.log('#orig_nodeCreated', this)
const widget = {
type: 'div',
name: 'result',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(this.div.style, {
...get_position_style(ctx, widget_width, y, node.size[1]),
flexWrap: 'wrap',
justifyContent: 'space-between',
// outline: '1px solid red',
paddingLeft: '0px',
width: widget_width + 'px'
})
}
}
widget.div = $el('div', {})
widget.div.className = 'prompt_image_output'
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
initLightBox()
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
const onResize = this.onResize
this.onResize = function () {
// 缩放发生
// console.log('##缩放发生', this.size)
let w = this.size[0] * 0.5 - 12
Array.from(widget.div.querySelectorAll('.card'), card => {
card.style.width = `${w}px`
card.style.height = `${
w * parseFloat(card.getAttribute('data-scale'))
}px`
})
return onResize?.apply(this, arguments)
}
// this.serialize_widgets = true //需要保存参数
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = async function (message) {
onExecuted?.apply(this, arguments)
console.log('#PromptImage', message.prompts, message._images)
// window._mixlab_app_json = message.json
try {
let widget = this.widgets.filter(w => w.name === 'result')[0]
widget.value = message
_createResult(this, widget, { ...message })
} catch (error) {
console.log(error)
}
}
this.serialize_widgets = true //需要保存参数
}
},
async loadedGraphNode (node, app) {
if (node.type === 'PromptImage') {
// await sleep(0)
let widget = node.widgets.filter(w => w.name === 'result')[0]
console.log('widget.value', widget.value)
initLightBox()
let cards = widget.div.querySelectorAll('.card')
if (cards.length == 0) node.size = [280, 120]
if(widget.value) _createResult(node, widget, widget.value)
}
}
})
-203
View File
@@ -1,203 +0,0 @@
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: `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)
// }
// }
// }
}
})
+110
View File
@@ -0,0 +1,110 @@
import { app } from '../../../scripts/app.js'
import { $el } from '../../../scripts/ui.js'
import { api } from '../../../scripts/api.js'
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: 'row',
// alignItems: 'center',
justifyContent: 'flex-start'
}
}
app.registerExtension({
name: 'Mixlab.share.ShareToWeibo',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'ShareToWeibo') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
// console.log(this)
const widget = {
type: 'div',
name: 'ShareToWeiboBtn',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(
ctx,
widget_width,
node.widgets[2].last_y +16,
node.size[1]
)
)
}
}
const style = `
flex-direction: row;
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;
color: var(--descrip-text);`
widget.div = $el('div', {})
const btn = document.createElement('button')
btn.innerText = 'Share'
btn.style = style
btn.addEventListener('click', () => {
if (window._mixlab_share_to_weibo)
window.open(window._mixlab_share_to_weibo)
})
document.body.appendChild(widget.div)
widget.div.appendChild(btn)
this.addCustomWidget(widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
this.serialize_widgets = true //需要保存参数
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = async function (message) {
onExecuted?.apply(this, arguments)
// console.log(this.widgets)
window._mixlab_share_to_weibo = message.url
try {
const div = this.widgets.filter(w => w.div)[0].div
Array.from(
div.querySelectorAll('button'),
b => (b.style.background = 'yellow')
)
} catch (error) {}
}
}
}
})
-329
View File
@@ -1,329 +0,0 @@
const smart_connect_config_input = [
{
node_type: 'CLIPTextEncode',
node_widget_name: 'text',
inputNodeName: 'RandomPrompt',
inputNode_output_name: 'STRING'
},
{
node_type: 'CLIPTextEncode',
node_widget_name: 'text',
inputNodeName: 'EmbeddingPrompt',
inputNode_output_name: 'STRING'
},
{
node_type: 'CLIPTextEncode',
node_widget_name: 'text',
inputNodeName: 'ChinesePrompt_Mix',
inputNode_output_name: 'prompt'
},
{
node_type: 'CheckpointLoaderSimple',
node_widget_name: 'ckpt_name',
inputNodeName: 'CkptNames_',
inputNode_output_name: 'ckpt_names'
},
{
node_type: 'KSampler',
node_widget_name: 'sampler_name',
inputNodeName: 'SamplerNames_',
inputNode_output_name: 'sampler_names'
},
{
node_type: 'LoraLoaderModelOnly',
node_widget_name: 'lora_name',
inputNodeName: 'LoraNames_',
inputNode_output_name: 'lora_names'
},
{
node_type: 'LoadLoRA',
node_widget_name: 'lora_name',
inputNodeName: 'LoraNames_',
inputNode_output_name: 'lora_names'
},
{
node_type: 'Moondream',
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'
}
]
const smart_connect_config_output = [
{
node_type: 'LoadImage',
node_output_name: 'IMAGE',
outputNodeName: 'ClipInterrogator',
outputNode_input_name: 'image'
},
{
node_type: 'VAEDecode',
node_output_name: 'IMAGE',
outputNodeName: 'PromptImage',
outputNode_input_name: 'images'
},
{
node_type: 'VAEDecode',
node_output_name: 'IMAGE',
outputNodeName: 'PreviewImage',
outputNode_input_name: 'images'
},
{
node_type: 'VAEDecode',
node_output_name: 'IMAGE',
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',
outputNodeName: 'ShowTextForGPT',
outputNode_input_name: 'text'
}
]
// import {
// convertToInput,
// getConfig,
// isConvertableWidget
// } from '../../../extensions/core/widgetInputs.js'
const CONVERTED_TYPE = 'converted-widget'
const GET_CONFIG = Symbol()
function getConfig (widgetName) {
const { nodeData } = this.constructor
return (
nodeData?.input?.required[widgetName] ??
nodeData?.input?.optional?.[widgetName]
)
}
function hideWidget (node, widget, suffix = '') {
widget.origType = widget.type
widget.origComputeSize = widget.computeSize
widget.origSerializeValue = widget.serializeValue
widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically
widget.type = CONVERTED_TYPE + suffix
widget.serializeValue = () => {
// Prevent serializing the widget if we have no input linked
if (!node.inputs) {
return undefined
}
let node_input = node.inputs.find(i => i.widget?.name === widget.name)
if (!node_input || !node_input.link) {
return undefined
}
return widget.origSerializeValue
? widget.origSerializeValue()
: widget.value
}
// Hide any linked widgets, e.g. seed+seedControl
if (widget.linkedWidgets) {
for (const w of widget.linkedWidgets) {
hideWidget(node, w, ':' + widget.name)
}
}
}
function convertToInput (node, widget, config) {
hideWidget(node, widget)
const type = config[0]
// Add input and store widget config for creating on primitive node
const sz = node.size
node.addInput(widget.name, type, {
widget: { name: widget.name, [GET_CONFIG]: () => config }
})
for (const widget of node.widgets) {
widget.last_y += LiteGraph.NODE_SLOT_HEIGHT
}
// Restore original size but grow if needed
node.setSize([Math.max(sz[0], node.size[0]), Math.max(sz[1], node.size[1])])
}
export function smart_init () {
LGraphCanvas.prototype._createNodeForInput = function (
node,
widget,
inputNodeName,
inputNode_slot
) {
// console.log(node.pos)
// var widget = node.widgets.filter(w => w.name === node_widget_name)[0]
if (widget) {
// 如果有存在的,没有连线输出的,自动连,不新建
let input_node = null
Array.from(app.graph.findNodesByType(inputNodeName), n => {
var links = n.outputs.filter(o => o.name === inputNode_slot)[0].links
// console.log(links)
if (!links || links?.length === 0) input_node = n
})
// 新建
if (!input_node) {
input_node = LiteGraph.createNode(inputNodeName)
input_node.pos = [node.pos[0] - node.size[0] - 24, node.pos[1] - 48]
app.canvas.graph.add(input_node, false)
} else {
input_node.pos = [node.pos[0] - node.size[0] - 24, node.pos[1] - 48]
}
const config = getConfig.call(node, widget.name) ?? [
widget.type,
widget.options || {}
]
let node_slotType = config[0]
// 如果input没有,则创建
if (
!node.inputs?.filter(inp => inp.name === widget.name)[0] ||
!node.inputs
)
convertToInput(node, widget, config)
input_node.connectByType(inputNode_slot, node, node_slotType)
}
}
LGraphCanvas.prototype._createNodeForOutput = function (
node,
widget,
outputNodeName,
outputNode_slot
) {
if (widget) {
let output_node
Array.from(app.graph.findNodesByType(outputNodeName), n => {
var links = n.inputs.filter(o => o.name === outputNode_slot)[0].links
// console.log(links)
if (!links || links?.length === 0) output_node = n
})
console.log('output_node', output_node, widget.name)
if (!output_node) {
// 新建
output_node = LiteGraph.createNode(outputNodeName)
output_node.pos = [node.pos[0] + node.size[0] + 24, node.pos[1] - 48]
app.canvas.graph.add(output_node, false)
} else {
output_node.pos = [node.pos[0] + node.size[0] + 24, node.pos[1] - 48]
}
const config = getConfig.call(node, widget.name) ?? [
widget.type,
widget.options || {}
]
let node_slotType = config[0]
console.log(node_slotType, output_node, outputNode_slot)
let type = output_node.inputs.filter(
inp => inp.name == outputNode_slot
)[0].type
node.connectByType(node_slotType, output_node, type)
}
}
}
export function addSmartMenu (options, node) {
let sopts = []
for (const sc of smart_connect_config_input) {
// 有智能推荐,则出现
if (node.type === sc.node_type) {
// console.log('smart',node)
// 则出现 randomPrompt
// CLIPTextEncode 的widget ,name== 'text'
let node_widget_name = sc.node_widget_name
let widget = node.widgets.filter(w => w.name === node_widget_name)[0]
if (!widget) {
// 控件没有,则查找inputs
widget = node.inputs.filter(w => w.name === node_widget_name)[0]
}
let isLinkNull = true
// 如果input里已经有,但是link为空
if (node.inputs?.filter(inp => inp.name === node_widget_name)[0]) {
isLinkNull =
node.inputs.filter(inp => inp.name === node_widget_name)[0].link ===
null
}
if (widget && isLinkNull) {
sopts.push({
content: sc.inputNodeName.split('_')[0] + '➡️',
callback: () => {
LGraphCanvas.prototype._createNodeForInput(
node, //当前node
widget, //当前node里需要自动连线的widget
sc.inputNodeName, //作为input的node type
sc.inputNode_output_name // 作为input的node的outputs的name. the input slot type of the target node
)
}
})
}
}
}
for (const sc of smart_connect_config_output) {
if (node.type === sc.node_type) {
let node_output_name = sc.node_output_name
const widget = node.outputs.filter(w => w.name === node_output_name)[0]
let isLinkNull = true
// 如果output里 link为空
if (node.outputs?.filter(inp => inp.name === node_output_name)[0]) {
isLinkNull =
node.outputs.filter(inp => inp.name === node_output_name)[0].links
?.length === 0
if (!node.outputs.filter(inp => inp.name === node_output_name)[0].links)
isLinkNull = true
}
if (widget && isLinkNull) {
sopts.push({
content: '➡️' + sc.outputNodeName.split('_')[0],
callback: () => {
LGraphCanvas.prototype._createNodeForOutput(
node, //当前node
widget, //当前node里需要自动连线的widget
sc.outputNodeName, //作为input的node type
sc.outputNode_input_name // 作为input的node的outputs的name. the input slot type of the target node
)
}
})
}
}
}
if (sopts.length > 0) options = [...sopts, null, ...options]
return options
}
-295
View File
@@ -1,295 +0,0 @@
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: `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
@@ -1,323 +0,0 @@
// 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()
+205 -1748
View File
File diff suppressed because it is too large Load Diff

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