Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
38851372e1 | ||
|
|
1e9ffc5ffc | ||
|
|
ccb17f18b8 | ||
|
|
380c596d9a | ||
|
|
4c3328797b | ||
|
|
d22f1f44f8 | ||
|
|
fb435d47ea | ||
|
|
d9d597bf83 | ||
|
|
1913c65d6f | ||
|
|
b8b24040eb | ||
|
|
e9bed88d63 | ||
|
|
d95772147f | ||
|
|
7cbc2a2a2b | ||
|
|
63811206c9 | ||
|
|
48a4b5bfc6 | ||
|
|
7980f4eef4 | ||
|
|
b19c8b07ea | ||
|
|
3f506d265c | ||
|
|
d8af918e7d | ||
|
|
33566e8474 | ||
|
|
ea0350c2dc | ||
|
|
24526623fb | ||
|
|
7715ebfd06 | ||
|
|
2ba814c131 | ||
|
|
d5ad332666 | ||
|
|
fee7e7bf73 | ||
|
|
01f17ff02b | ||
|
|
da7120219a | ||
|
|
d616d18069 | ||
|
|
dbf76f288c | ||
|
|
05124006ba | ||
|
|
4586af311c | ||
|
|
1cea58c7cf | ||
|
|
a84f7c4a58 | ||
|
|
1d9bf86560 | ||
|
|
8d352b85bc | ||
|
|
403562575c | ||
|
|
dbe2cd6569 | ||
|
|
b3b0a961c5 | ||
|
|
9dfb8b9c15 | ||
|
|
b4ea58946c | ||
|
|
513fc4b67e | ||
|
|
924a16e31c | ||
|
|
ad43f8e0bb | ||
|
|
b7e1ce8a3c | ||
|
|
299090d184 | ||
|
|
3aa7bca86b | ||
|
|
56de6f0bc9 | ||
|
|
f6b5f5c99c | ||
|
|
79b81100ff | ||
|
|
c885fd9bcf | ||
|
|
c46aaa6084 | ||
|
|
0562d4eb0e | ||
|
|
cfa56d36d7 | ||
|
|
913cfe73ae | ||
|
|
590161c560 | ||
|
|
d674313240 | ||
|
|
44258090cd | ||
|
|
1222799e8f | ||
|
|
ecc972f7b4 | ||
|
|
c079652878 | ||
|
|
b333bf05c6 | ||
|
|
f52a8fd40b | ||
|
|
f320647d78 | ||
|
|
631acfe44a | ||
|
|
7269f02f8a | ||
|
|
fdc761ebfa | ||
|
|
689d988130 | ||
|
|
e29fd5ed24 | ||
|
|
35f19b75fa | ||
|
|
96125a65a0 | ||
|
|
70ca7cc35f | ||
|
|
5858fb0606 | ||
|
|
149fab2105 | ||
|
|
39c5ccf469 | ||
|
|
cd5fcd1f70 | ||
|
|
96b3897f95 | ||
|
|
293692398e | ||
|
|
84307b95aa | ||
|
|
1e1b82399c | ||
|
|
8c610dd8f9 | ||
|
|
ac7db29df9 | ||
|
|
d5fc43e11f | ||
|
|
65ba451d5d | ||
|
|
cc79b11c12 | ||
|
|
148b9a695a | ||
|
|
83ca2f8ab6 | ||
|
|
514f93c892 | ||
|
|
9d64ea448b | ||
|
|
1ffecfcc7d | ||
|
|
aadefaf40b | ||
|
|
4c25580295 | ||
|
|
192181007a | ||
|
|
54ea3af9b8 | ||
|
|
c0ca3c7e95 | ||
|
|
5c8af8f0b4 |
@@ -0,0 +1,21 @@
|
||||
name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
+5
-1
@@ -1,4 +1,6 @@
|
||||
__pycache__
|
||||
.DS_Store
|
||||
*.cache
|
||||
*.ini
|
||||
*.bak
|
||||
wildcards/**
|
||||
@@ -8,4 +10,6 @@ autocomplete/**
|
||||
docs/**
|
||||
.vscode/
|
||||
.idea/
|
||||
mmb-preset.custom.txt
|
||||
mmb-preset.custom.txt
|
||||
config.yaml
|
||||
node.tar.gz
|
||||
+140
-56
@@ -9,43 +9,101 @@
|
||||
|
||||
**ComfyUI-Easy-Use** is a simplified node integration package, which is extended on the basis of [tinyterraNodes](https://github.com/TinyTerra/ComfyUI_tinyterraNodes), and has been integrated and optimized for many mainstream node packages to achieve the purpose of faster and more convenient use of ComfyUI. While ensuring the degree of freedom, it restores the ultimate smooth image production experience that belongs to Stable Diffusion.
|
||||
|
||||
## Introduce
|
||||
|
||||
### Random seed control before generate
|
||||
<img src="https://raw.githubusercontent.com/yolain/yolain-comfyui-workflow/main/Docs/seed_generate_compare.jpg">
|
||||
|
||||
### Separate sampling parameters from sample preview
|
||||
|
||||
<img src="https://raw.githubusercontent.com/yolain/yolain-comfyui-workflow/main/Docs/workflow_node_compare.png">
|
||||
|
||||
### Wildcard prompt nodes are supported
|
||||
## Introduce
|
||||
|
||||
Support `.yaml`, `.txt`, `.json` format wildcard files, just place them in the 'wildcards' folder of the node package, and update the file to run ComfyUI again. <br>
|
||||
To use the Lora Block Weight usage, make sure that [ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) is installed in the custom node package.
|
||||
|
||||
### UI interface is beautified
|
||||
|
||||
After installing the node package, the UI interface will be automatically switched, if you need to change other themes, please switch and refresh the page in Settings -> Color Palette.
|
||||
|
||||
### Stable Cascade
|
||||
|
||||
[WorkFlow Example](https://github.com/yolain/ComfyUI-Easy-Use/blob/main/README.en.md#StableCascade) <br><br>
|
||||
|
||||
Usage:<br>
|
||||
1.There is no need to load the rest of the VAE and clips when you are choose [checkpoints](https://huggingface.co/stabilityai/stable-cascade/tree/main/comfyui_checkpoints) models.<br>
|
||||
2.You need to load it extra [stage_a](https://huggingface.co/stabilityai/stable-cascade/blob/main/stage_a.safetensors)、[clip](https://huggingface.co/stabilityai/stable-cascade/resolve/main/text_encoder/model.safetensors) and [effnet_encoder](https://huggingface.co/stabilityai/stable-cascade/resolve/main/effnet_encoder.safetensors?download=true)、[previewer](https://huggingface.co/stabilityai/stable-cascade/resolve/main/previewer.safetensors) for img2img when you are choose unet models.<br>
|
||||
<br>
|
||||
|
||||
### Layer Diffusion
|
||||
|
||||
[WorkFlow Example](https://github.com/yolain/ComfyUI-Easy-Use/blob/main/README.en.md#LayerDiffusion) <br><br>
|
||||
|
||||
Usage:<br>
|
||||
you need to run `pip install -r requirements.txt` to install python dependencies when **diffusers** was not installed.
|
||||
- Inspire by [tinyterraNodes](https://github.com/TinyTerra/ComfyUI_tinyterraNodes), which greatly reduces the time cost of tossing workflows。
|
||||
- UI interface beautification, the first time you install the user, if you need to use the UI theme, please switch the theme in Settings -> Color Palette and refresh page.
|
||||
- Added a node for pre-sampling parameter configuration, which can be separated from the sampling node for easier previewing
|
||||
- Wildcards and lora's are supported, for Lora Block Weight usage, ensure that the custom node package has the [ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack)
|
||||
- Multi-selectable styled cue word selector, default is Fooocus style json, custom json can be placed under styles, samples folder can be placed in the preview image (name and name consistent, image file name such as spaces need to be converted to underscores '_')
|
||||
- The loader enables the A1111 prompt mode, which reproduces nearly identical images to those generated by webui, and needs to be installed [ComfyUI_smZNodes](https://github.com/shiimizu/ComfyUI_smZNodes) first.
|
||||
- Noise injection into the latent space can be achieved using the `easy latentNoisy` or `easy preSamplingNoiseIn` node
|
||||
- Simplified processes for SD1.x, SD2.x, SDXL, SVD, Zero123, etc. [Example](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#StableDiffusion)
|
||||
- Simplified Stable Cascade [Example](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#StableCascade)
|
||||
- Simplified Layer Diffuse [Example](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#LayerDiffusion),The first time you use it you may need to run `pip install -r requirements.txt` to install the required dependencies.
|
||||
- Simplified InstantID [Example](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#InstantID), You need to make sure that the custom node package has the [ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID)
|
||||
- Extending the usability of XYplot
|
||||
- Fooocus Inpaint integration
|
||||
- Integration of common logical calculations, conversion of types, display of all types, etc.
|
||||
- Background removal nodes for the RMBG-1.4 model supporting BriaAI, [BriaAI Guide](https://huggingface.co/briaai/RMBG-1.4)
|
||||
- Forcibly cleared the memory usage of the comfy UI model are supported
|
||||
- Stable Diffusion 3 multi-account API nodes are supported
|
||||
|
||||
## Changelog
|
||||
|
||||
**v1.1.1 (2024/3/16)**
|
||||
**v1.1.8**
|
||||
|
||||
- Added `easy controlnetStack`
|
||||
- Added `easy applyBrushNet` - [Workflow Example](https://github.com/yolain/ComfyUI-Yolain-Workflows/blob/main/workflows/2_advanced/2-4inpainting/2-4brushnet_1.1.8.json)
|
||||
- Added `easy applyPowerPaint` - [Workflow Example](https://github.com/yolain/ComfyUI-Yolain-Workflows/blob/main/workflows/2_advanced/2-4inpainting/2-4powerpaint_outpaint_1.1.8.json)
|
||||
|
||||
**v1.1.7**
|
||||
|
||||
- Added `easy prompt` - Subject and light presets, maybe adjusted later
|
||||
- Added `easy icLightApply` - Light and shadow migration, Code based on [ComfyUI-IC-Light](https://github.com/huchenlei/ComfyUI-IC-Light)
|
||||
- Added `easy imageSplitGrid`
|
||||
- `easy kSamplerInpainting` added options such as different diffusion and brushnet in **additional** widget
|
||||
- Support for brushnet model loading - [ComfyUI-BrushNet](https://github.com/nullquant/ComfyUI-BrushNet)
|
||||
- Added `easy applyFooocusInpaint` - Replace FooocusInpaintLoader
|
||||
- Removed `easy fooocusInpaintLoader`
|
||||
|
||||
**v1.1.6**
|
||||
|
||||
- Added **alignYourSteps** to **schedulder** widget in all `easy preSampling` and `easy fullkSampler`
|
||||
- Added **Preview&Choose** to **image_output** widget in `easy kSampler` & `easy fullkSampler`
|
||||
- Added `easy styleAlignedBatchAlign` - Credit of [style_aligned_comfy](https://github.com/brianfitzgerald/style_aligned_comfy)
|
||||
- Added `easy ckptNames`
|
||||
- Added `easy controlnetNames`
|
||||
- Added `easy imagesSplitimage` - Batch images split into single images
|
||||
- Added `easy imageCount` - Get Image Count
|
||||
- Added `easy textSwitch` - Text Switch
|
||||
|
||||
**v1.1.5**
|
||||
|
||||
- Rewrite `easy cleanGPUUsed` - the memory usage of the comfyUI can to be cleared
|
||||
- Added `easy humanSegmentation` - Human Part Segmentation
|
||||
- Added `easy imageColorMatch`
|
||||
- Added `easy ipadapterApplyRegional`
|
||||
- Added `easy ipadapterApplyFromParams`
|
||||
- Added `easy imageInterrogator` - Image To Prompt
|
||||
- Added `easy stableDiffusion3API` - Easy Stable Diffusion 3 Multiple accounts API Node
|
||||
|
||||
**v1.1.4**
|
||||
|
||||
- Added `easy preSamplingCustom` - Custom-PreSampling, can be supported cosXL-edit
|
||||
- Added `easy ipadapterStyleComposition`
|
||||
- Added the right-click menu to view checkpoints and lora information in all Loaders
|
||||
- Fixed `easy preSamplingNoiseIn`、`easy latentNoisy`、`east Unsampler` compatible with ComfyUI Revision>=2098 [0542088e] or later
|
||||
|
||||
|
||||
**v1.1.3**
|
||||
|
||||
- `easy ipadapterApply` Added **COMPOSITION** preset
|
||||
- Supported [ResAdapter](https://huggingface.co/jiaxiangc/res-adapter) when load ResAdapter lora
|
||||
- Added `easy promptLine`
|
||||
- Added `easy promptReplace`
|
||||
- Added `easy promptConcat`
|
||||
- `easy wildcards` Added **multiline_mode**
|
||||
|
||||
**v1.1.2**
|
||||
|
||||
- Optimized some of the recommended nodes for slots related to EasyUse
|
||||
- Added **Enable ContextMenu Auto Nest Subdirectories** The setting item is enabled by default, and it can be classified into subdirectories, checkpoints and loras previews
|
||||
- Added `easy sv3dLoader`
|
||||
- Added `easy dynamiCrafterLoader`
|
||||
- Added `easy ipadapterApply`
|
||||
- Added `easy ipadapterApplyADV`
|
||||
- Added `easy ipadapterApplyEncoder`
|
||||
- Added `easy ipadapterApplyEmbeds`
|
||||
- Added `easy preMaskDetailerFix`
|
||||
- Fixed `easy stylesSelector` is change the prompt when not select the style
|
||||
- Fixed `easy pipeEdit` error when add lora to prompt
|
||||
- Fixed layerDiffuse xyplot bug
|
||||
- `easy kSamplerInpainting` add *additional* widget,you can choose 'Differential Diffusion' or 'Only InpaintModelConditioning'
|
||||
|
||||
**v1.1.1**
|
||||
|
||||
- The issue that the seed is 0 when a node with a seed control is added and **control before generate** is fixed for the first time run queue prompt.
|
||||
- `easy preSamplingAdvanced` Added **return_with_leftover_noise**
|
||||
@@ -55,7 +113,7 @@ you need to run `pip install -r requirements.txt` to install python dependencies
|
||||
- Remove forced **control_before_generate** settings。 If you want to use control_before_generate, change widget_value_control_mode to before in system settings
|
||||
- Added `easy imageRemBg` - The default is BriaAI's RMBG-1.4 model, which removes the background effect more and faster
|
||||
|
||||
**v1.1.0 (d5ff84e)**
|
||||
**v1.1.0**
|
||||
|
||||
- Added `easy imageSplitList` - to split every N images
|
||||
- Added `easy preSamplingDiffusionADDTL` - It can modify foreground、background or blended additional prompt
|
||||
@@ -71,15 +129,18 @@ you need to run `pip install -r requirements.txt` to install python dependencies
|
||||
- Fixed `easy instantIDApply` mask not input right
|
||||
|
||||
|
||||
**v1.0.9 (ff1add1)**
|
||||
<details>
|
||||
<summary><b>v1.0.9</b></summary>
|
||||
|
||||
- Fixed the error when ComfyUI-Impack-Pack and ComfyUI_InstantID were not installed
|
||||
- Fixed `easy pipeIn`
|
||||
- Added `easy instantIDApply` - you need installed [ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID) fisrt, Workflow[Example](https://github.com/yolain/ComfyUI-Easy-Use/blob/main/README.en.md#InstantID)
|
||||
- Fixed `easy detailerFix` not added to the list of nodes available for saving images formatting extensions
|
||||
- Fixed `easy XYInputs: PromptSR` errors are reported when replacing negative prompts
|
||||
</details>
|
||||
|
||||
**v1.0.8 (f28cbf7)**
|
||||
<details>
|
||||
<summary><b>v1.0.8</b></summary>
|
||||
|
||||
- `easy cascadeLoader` stage_c and stage_b support the checkpoint model (Download [checkpoints](https://huggingface.co/stabilityai/stable-cascade/tree/main/comfyui_checkpoints) models)
|
||||
- `easy styleSelector` The search box is modified to be case-insensitive
|
||||
@@ -89,8 +150,10 @@ you need to run `pip install -r requirements.txt` to install python dependencies
|
||||
- Fixed the error of SDXLClipModel in ComfyUI revision 2016[c2cb8e88] and above (the revision number was judged to be compatible with the old revision)
|
||||
- Fixed `easy detailerFix` generation error when batch size is greater than 1
|
||||
- Optimize the code, reduce a lot of redundant code and improve the running speed
|
||||
</details>
|
||||
|
||||
**v1.0.7 (2024-02-19)**
|
||||
<details>
|
||||
<summary><b>v1.0.7</b></summary>
|
||||
|
||||
- Added `easy cascadeLoader` - stable cascade Loader
|
||||
- Added `easy preSamplingCascade` - stable cascade preSampling Settings
|
||||
@@ -98,22 +161,28 @@ you need to run `pip install -r requirements.txt` to install python dependencies
|
||||
- Added `easy cascadeKSampler` - stable cascade stage-c ksampler simple
|
||||
-
|
||||
- Optimize the image to image[Example](https://github.com/yolain/ComfyUI-Easy-Use/blob/main/README.en.md#image-to-image)
|
||||
</details>
|
||||
|
||||
**v1.0.6**
|
||||
<details>
|
||||
<summary><b>v1.0.6</b></summary>
|
||||
|
||||
- Added `easy XYInputs: Checkpoint`
|
||||
- Added `easy XYInputs: Lora`
|
||||
- `easy seed` can manually switch the random seed when increasing the fixed seed value
|
||||
- Fixed `easy fullLoader` and all loaders to automatically adjust the node size when switching LoRa
|
||||
- Removed the original ttn image saving logic and adapted to the default image saving format extension of ComfyUI
|
||||
</details>
|
||||
|
||||
- **v1.0.5**
|
||||
<details>
|
||||
<summary><b>v1.0.5</b></summary>
|
||||
|
||||
- Added `easy isSDXL`
|
||||
- Added prompt word control on `easy svdLoader`, which can be used with open_clip model
|
||||
- Added **populated_text** on `easy wildcards`, wildcard populated text can be output
|
||||
</details>
|
||||
|
||||
**v1.0.4**
|
||||
<details>
|
||||
<summary><b>v1.0.4</b></summary>
|
||||
|
||||
- `easy showAnything` added support for converting other types (e.g., tensor conditions, images, etc.)
|
||||
- Added `easy showLoaderSettingsNames` can display the model and VAE name in the output loader assembly
|
||||
@@ -133,9 +202,10 @@ you need to run `pip install -r requirements.txt` to install python dependencies
|
||||
- Changing the first-time install node package no longer automatically replaces the theme, you need to manually adjust and refresh the page
|
||||
- `easy imageSave` added **only_preivew**
|
||||
- Adjust the `easy latentCompositeMaskedWithCond` node
|
||||
</details>
|
||||
|
||||
|
||||
**v1.0.3**
|
||||
<details>
|
||||
<summary><b>v1.0.3</b></summary>
|
||||
|
||||
- Added `easy stylesSelector`
|
||||
- Added **scale_soft_weights** in `easy controlnetLoader` and `easy controlnetLoaderADV`
|
||||
@@ -153,9 +223,10 @@ you need to run `pip install -r requirements.txt` to install python dependencies
|
||||
|
||||
- Adjust the UI theme, divided into two sets of styles: the official default background and the dark black background, which can be switched in the color palette in the settings
|
||||
- Modify the styles path to be compatible with other environments
|
||||
</details>
|
||||
|
||||
|
||||
**v1.0.2**
|
||||
<details>
|
||||
<summary><b>v1.0.2</b></summary>
|
||||
|
||||
- Added `easy XYPlotAdvanced` and some nodes about `easy XYInputs`
|
||||
- Added **Alt+1-Alt+9** Shortcut keys to quickly paste node presets for Node templates (corresponding to 1~9 sequences)
|
||||
@@ -172,7 +243,7 @@ you need to run `pip install -r requirements.txt` to install python dependencies
|
||||
- Removed `easy imageRemBg`
|
||||
- Remove the introductory diagram and workflow files from the package to reduce the package size
|
||||
- Replaced the font file used in the generation of XY diagrams
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>v1.0.1</b></summary>
|
||||
@@ -252,19 +323,22 @@ you need to run `pip install -r requirements.txt` to install python dependencies
|
||||
|
||||
Disclaimer: Opened source was not easy. I have a lot of respect for the contributions of these original authors. I just did some integration and optimization.
|
||||
|
||||
| Nodes Name(Search Name) | Related libraries | Library-related node |
|
||||
|:---------------------------|:----------------------------------------------------------------------------|:-------------------------|
|
||||
| easy setNode | [ComfyUI-extensions](https://github.com/diffus3/ComfyUI-extensions) | diffus3.SetNode |
|
||||
| easy getNode | [ComfyUI-extensions](https://github.com/diffus3/ComfyUI-extensions) | diffus3.GetNode |
|
||||
| easy bookmark | [rgthree-comfy](https://github.com/rgthree/rgthree-comfy) | Bookmark 🔖 |
|
||||
| easy portraitMarker | [comfyui-portrait-master](https://github.com/florestefano1975/comfyui-portrait-master) | Portrait Master |
|
||||
| easy LLLiteLoader | [ControlNet-LLLite-ComfyUI](https://github.com/kohya-ss/ControlNet-LLLite-ComfyUI) | LLLiteLoader |
|
||||
| easy globalSeed | [ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) | Global Seed (Inspire) |
|
||||
| easy preSamplingDynamicCFG | [sd-dynamic-thresholding](https://github.com/mcmonkeyprojects/sd-dynamic-thresholding) | DynamicThresholdingFull |
|
||||
| dynamicThresholdingFull | [sd-dynamic-thresholding](https://github.com/mcmonkeyprojects/sd-dynamic-thresholding) | DynamicThresholdingFull |
|
||||
| easy imageInsetCrop | [rgthree-comfy](https://github.com/rgthree/rgthree-comfy) | ImageInsetCrop |
|
||||
| easy poseEditor | [ComfyUI_Custom_Nodes_AlekPet](https://github.com/AlekPet/ComfyUI_Custom_Nodes_AlekPet) | poseNode |
|
||||
| Nodes Name(Search Name) | Related libraries | Library-related node |
|
||||
|:-------------------------------|:----------------------------------------------------------------------------|:-------------------------|
|
||||
| easy setNode | [ComfyUI-extensions](https://github.com/diffus3/ComfyUI-extensions) | diffus3.SetNode |
|
||||
| easy getNode | [ComfyUI-extensions](https://github.com/diffus3/ComfyUI-extensions) | diffus3.GetNode |
|
||||
| easy bookmark | [rgthree-comfy](https://github.com/rgthree/rgthree-comfy) | Bookmark 🔖 |
|
||||
| easy portraitMarker | [comfyui-portrait-master](https://github.com/florestefano1975/comfyui-portrait-master) | Portrait Master |
|
||||
| easy LLLiteLoader | [ControlNet-LLLite-ComfyUI](https://github.com/kohya-ss/ControlNet-LLLite-ComfyUI) | LLLiteLoader |
|
||||
| easy globalSeed | [ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) | Global Seed (Inspire) |
|
||||
| easy preSamplingDynamicCFG | [sd-dynamic-thresholding](https://github.com/mcmonkeyprojects/sd-dynamic-thresholding) | DynamicThresholdingFull |
|
||||
| dynamicThresholdingFull | [sd-dynamic-thresholding](https://github.com/mcmonkeyprojects/sd-dynamic-thresholding) | DynamicThresholdingFull |
|
||||
| easy imageInsetCrop | [rgthree-comfy](https://github.com/rgthree/rgthree-comfy) | ImageInsetCrop |
|
||||
| easy poseEditor | [ComfyUI_Custom_Nodes_AlekPet](https://github.com/AlekPet/ComfyUI_Custom_Nodes_AlekPet) | poseNode |
|
||||
| easy preSamplingLayerDiffusion | [ComfyUI-layerdiffusion](https://github.com/huchenlei/ComfyUI-layerdiffusion) | LayeredDiffusionApply... |
|
||||
| easy dynamiCrafterLoader | [ComfyUI-layerdiffusion](https://github.com/ExponentialML/ComfyUI_Native_DynamiCrafter) | Apply Dynamicrafter |
|
||||
| easy imageChooser | [cg-image-picker](https://github.com/chrisgoringe/cg-image-picker) | Preview Chooser |
|
||||
| easy styleAlignedBatchAlign | [style_aligned_comfy](https://github.com/chrisgoringe/cg-image-picker) | styleAlignedBatchAlign |
|
||||
|
||||
## Workflow Examples
|
||||
|
||||
@@ -306,4 +380,14 @@ Disclaimer: Opened source was not easy. I have a lot of respect for the contribu
|
||||
|
||||
[ComfyUI-Impact-Pack](https://github.com/ltdrdata/ComfyUI-Impact-Pack) - General modpack 1
|
||||
|
||||
[ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) - General Modpack 2
|
||||
[ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) - General Modpack 2
|
||||
|
||||
[ComfyUI-ResAdapter](https://github.com/jiaxiangc/ComfyUI-ResAdapter) - Make model generation independent of training resolution
|
||||
|
||||
[ComfyUI_IPAdapter_plus](https://github.com/cubiq/ComfyUI_IPAdapter_plus) - Style migration
|
||||
|
||||
[ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID) - Face migration
|
||||
|
||||
[ComfyUI-Custom-Scripts](https://github.com/pythongosssss/ComfyUI-Custom-Scripts) - pyssss🐍
|
||||
|
||||
[cg-image-picker](https://github.com/chrisgoringe/cg-image-picker) - Image Preview Chooser
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
|
||||
**ComfyUI-Easy-Use** 是一个化繁为简的节点整合包, 在 [tinyterraNodes](https://github.com/TinyTerra/ComfyUI_tinyterraNodes) 的基础上进行延展,并针对了诸多主流的节点包做了整合与优化,以达到更快更方便使用ComfyUI的目的,在保证自由度的同时还原了本属于Stable Diffusion的极致畅快出图体验。
|
||||
|
||||
<img src="https://raw.githubusercontent.com/yolain/yolain-comfyui-workflow/main/Docs/workflow_node_compare.png">
|
||||
[](https://github.com/yolain/ComfyUI-Yolain-Workflows)
|
||||
|
||||
## 特色介绍
|
||||
|
||||
@@ -23,18 +23,99 @@
|
||||
- 可多选的风格化提示词选择器,默认是Fooocus的样式json,可自定义json放在styles底下,samples文件夹里可放预览图(名称和name一致,图片文件名如有空格需转为下划线'_')
|
||||
- 加载器可开启A1111提示词风格模式,可重现与webui生成近乎相同的图像,需先安装 [ComfyUI_smZNodes](https://github.com/shiimizu/ComfyUI_smZNodes)
|
||||
- 可使用`easy latentNoisy`或`easy preSamplingNoiseIn`节点实现对潜空间的噪声注入
|
||||
- 简化 SD1.x、SD2.x、SDXL、SVD、Zero123等流程 [示例参考](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#StableDiffusion)
|
||||
- 简化 Stable Cascade [示例参考](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#StableCascade)
|
||||
- 简化 Layer Diffuse [示例参考](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#LayerDiffusion), 首次使用您可能需要运行 `pip install -r requirements.txt` 安装所需依赖
|
||||
- 简化 InstantID [示例参考](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#InstantID), 需先保证自定义节点包中安装了 [ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID)
|
||||
- 简化 SD1.x、SD2.x、SDXL、SVD、Zero123等流程
|
||||
- 简化 Stable Cascade [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows?tab=readme-ov-file#1-13-stable-cascade)
|
||||
- 简化 Layer Diffuse [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows?tab=readme-ov-file#2-3-layerdiffusion)
|
||||
- 简化 InstantID [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows?tab=readme-ov-file#2-2-instantid), 需先保证自定义节点包中安装了 [ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID)
|
||||
- 简化 IPAdapter, 需先保证自定义节点包中安装最新版v2的 [ComfyUI_IPAdapter_plus](https://github.com/cubiq/ComfyUI_IPAdapter_plus)
|
||||
- 扩展 XYplot 的可用性
|
||||
- 整合了Fooocus Inpaint功能
|
||||
- 整合了常用的逻辑计算、转换类型、展示所有类型等
|
||||
- 支持节点上checkpoint、lora模型子目录分类及预览图 (请在设置中开启上下文菜单嵌套子目录)
|
||||
- 支持BriaAI的RMBG-1.4模型的背景去除节点,[技术参考](https://huggingface.co/briaai/RMBG-1.4)
|
||||
- 支持 强制清理comfyUI模型显存占用
|
||||
- 支持Stable Diffusion 3 多账号API节点
|
||||
- 支持IC-Light的应用 [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows?tab=readme-ov-file#2-5-ic-light) | [代码整合来源](https://github.com/huchenlei/ComfyUI-IC-Light) | [技术参考](https://github.com/lllyasviel/IC-Light)
|
||||
- 中文提示词自动识别,使用[opus-mt-zh-en模型](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en)
|
||||
|
||||
## 更新日志
|
||||
|
||||
**v1.1.1 (2024/3/21)**
|
||||
**v1.1.8**
|
||||
|
||||
- 增加中文提示词自动翻译,使用[opus-mt-zh-en模型](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en), 默认已对wildcard、lora正则处理, 其他需要保留的中文,可使用`@你的提示词@`包裹 (若依赖安装完成后报错, 请重启),测算大约会占0.3GB显存
|
||||
- 增加 `easy controlnetStack` - controlnet堆
|
||||
- 增加 `easy applyBrushNet` - [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows/blob/main/workflows/2_advanced/2-4inpainting/2-4brushnet_1.1.8.json)
|
||||
- 增加 `easy applyPowerPaint` - [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows/blob/main/workflows/2_advanced/2-4inpainting/2-4powerpaint_outpaint_1.1.8.json)
|
||||
|
||||
**v1.1.7**
|
||||
|
||||
- 修复 一些模型(如controlnet模型等)未成功写入缓存,导致修改前置节点束参数(如提示词)需要二次载入模型的问题
|
||||
- 增加 `easy prompt` - 主体和光影预置项,后期可能会调整
|
||||
- 增加 `easy icLightApply` - 重绘光影, 从[ComfyUI-IC-Light](https://github.com/huchenlei/ComfyUI-IC-Light)优化
|
||||
- 增加 `easy imageSplitGrid` - 图像网格拆分
|
||||
- `easy kSamplerInpainting` 的 **additional** 属性增加差异扩散和brushnet等相关选项
|
||||
- 增加 brushnet模型加载的支持 - [ComfyUI-BrushNet](https://github.com/nullquant/ComfyUI-BrushNet)
|
||||
- 增加 `easy applyFooocusInpaint` - Fooocus内补节点 替代原有的 FooocusInpaintLoader
|
||||
- 移除 `easy fooocusInpaintLoader` - 容易bug,不再使用
|
||||
- 修改 easy kSampler等采样器中并联的model 不再替换输出中pipe里的model
|
||||
|
||||
**v1.1.6**
|
||||
|
||||
- 增加步调齐整适配 - 在所有的预采样和全采样器节点中的 调度器(schedulder) 增加了 **alignYourSteps** 选项
|
||||
- `easy kSampler` 和 `easy fullkSampler` 的 **image_output** 增加 **Preview&Choose**选项
|
||||
- 增加 `easy styleAlignedBatchAlign` - 风格对齐 [style_aligned_comfy](https://github.com/brianfitzgerald/style_aligned_comfy)
|
||||
- 增加 `easy ckptNames`
|
||||
- 增加 `easy controlnetNames`
|
||||
- 增加 `easy imagesSplitimage` - 批次图像拆分单张
|
||||
- 增加 `easy imageCount` - 图像数量
|
||||
- 增加 `easy textSwitch` - 文字切换
|
||||
|
||||
**v1.1.5**
|
||||
|
||||
- 重写 `easy cleanGPUUsed` - 可强制清理comfyUI的模型显存占用
|
||||
- 增加 `easy humanSegmentation` - 多类分割、人像分割
|
||||
- 增加 `easy imageColorMatch`
|
||||
- 增加 `easy ipadapterApplyRegional`
|
||||
- 增加 `easy ipadapterApplyFromParams`
|
||||
- 增加 `easy imageInterrogator` - 图像反推
|
||||
- 增加 `easy stableDiffusion3API` - 简易的Stable Diffusion 3 多账号API节点
|
||||
|
||||
**v1.1.4**
|
||||
|
||||
- 增加 `easy imageChooser` - 从[cg-image-picker](https://github.com/chrisgoringe/cg-image-picker)简化的图片选择器
|
||||
- 增加 `easy preSamplingCustom` - 自定义预采样,可支持cosXL-edit
|
||||
- 增加 `easy ipadapterStyleComposition`
|
||||
- 增加 在Loaders上右键菜单可查看 checkpoints、lora 信息
|
||||
- 修复 `easy preSamplingNoiseIn`、`easy latentNoisy`、`east Unsampler` 以兼容ComfyUI Revision>=2098 [0542088e] 以上版本
|
||||
- 修复 FooocusInpaint修改ModelPatcher计算权重引发的问题,理应在生成model后重置ModelPatcher为默认值
|
||||
|
||||
**v1.1.3**
|
||||
|
||||
- `easy ipadapterApply` 增加 **COMPOSITION** 预置项
|
||||
- 增加 对[ResAdapter](https://huggingface.co/jiaxiangc/res-adapter) lora模型 的加载支持
|
||||
- 增加 `easy promptLine`
|
||||
- 增加 `easy promptReplace`
|
||||
- 增加 `easy promptConcat`
|
||||
- `easy wildcards` 增加 **multiline_mode**属性
|
||||
- 增加 当节点需要下载模型时,若huggingface连接超时,会切换至镜像地址下载模型
|
||||
|
||||
**v1.1.2**
|
||||
|
||||
- 改写 EasyUse 相关节点的部分插槽推荐节点
|
||||
- 增加 **启用上下文菜单自动嵌套子目录** 设置项,默认为启用状态,可分类子目录及checkpoints、loras预览图
|
||||
- 增加 `easy sv3dLoader`
|
||||
- 增加 `easy dynamiCrafterLoader`
|
||||
- 增加 `easy ipadapterApply`
|
||||
- 增加 `easy ipadapterApplyADV`
|
||||
- 增加 `easy ipadapterApplyEncoder`
|
||||
- 增加 `easy ipadapterApplyEmbeds`
|
||||
- 增加 `easy preMaskDetailerFix`
|
||||
- `easy kSamplerInpainting` 增加 **additional** 属性,可设置成 Differential Diffusion 或 Only InpaintModelConditioning
|
||||
- 修复 `easy stylesSelector` 当未选择样式时,原有提示词发生了变化
|
||||
- 修复 `easy pipeEdit` 提示词输入lora时报错
|
||||
- 修复 layerDiffuse xyplot相关bug
|
||||
|
||||
**v1.1.1**
|
||||
|
||||
- 修复首次添加含seed的节点且当前模式为control_before_generate时,seed为0的问题
|
||||
- `easy preSamplingAdvanced` 增加 **return_with_leftover_noise**
|
||||
@@ -45,7 +126,7 @@
|
||||
- 去除强制**control_before_generate**设定
|
||||
- 增加 `easy imageRemBg` - 默认为BriaAI的RMBG-1.4模型, 移除背景效果更加,速度更快
|
||||
|
||||
**v1.1.0 (d5ff84e)**
|
||||
**v1.1.0**
|
||||
|
||||
- 增加 `easy imageSplitList` - 拆分每 N 张图像
|
||||
- 增加 `easy preSamplingDiffusionADDTL` - 可配置前景、背景、blended的additional_prompt等
|
||||
@@ -59,15 +140,18 @@
|
||||
- 修复 `easy instantIDApply` mask 未传入正确值
|
||||
- 修复 在 非a1111提示词风格下 BREAK 不生效的问题
|
||||
|
||||
**v1.0.9 (ff1add1)**
|
||||
<details>
|
||||
<summary><b>v1.0.9</b></summary>
|
||||
|
||||
- 修复未安装 ComfyUI-Impack-Pack 和 ComfyUI_InstantID 时报错
|
||||
- 修复 `easy pipeIn` - pipe设为可不必选
|
||||
- 增加 `easy instantIDApply` - 需要先安装 [ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID), 工作流参考[示例](https://github.com/yolain/ComfyUI-Easy-Use?tab=readme-ov-file#InstantID)
|
||||
- 增加 `easy instantIDApply` - 需要先安装 [ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID), 工作流参考[示例](https://github.com/yolain/ComfyUI-Yolain-Workflows?tab=readme-ov-file#2-2-instantid)
|
||||
- 修复 `easy detailerFix` 未添加到保存图片格式化扩展名可用节点列表
|
||||
- 修复 `easy XYInputs: PromptSR` 在替换负面提示词时报错
|
||||
</details>
|
||||
|
||||
**v1.0.8 (f28cbf7)**
|
||||
<details>
|
||||
<summary><b>v1.0.8</b></summary>
|
||||
|
||||
- `easy cascadeLoader` stage_c 与 stage_b 支持checkpoint模型 (需要下载[checkpoints](https://huggingface.co/stabilityai/stable-cascade/tree/main/comfyui_checkpoints))
|
||||
- `easy styleSelector` 搜索框修改为不区分大小写匹配
|
||||
@@ -81,13 +165,16 @@
|
||||
|
||||
(翻译对照已由 [AIGODLIKE-COMFYUI-TRANSLATION](https://github.com/AIGODLIKE/AIGODLIKE-ComfyUI-Translation) 统一维护啦!
|
||||
首次下载或者版本较早的朋友请更新 AIGODLIKE-COMFYUI-TRANSLATION 和本节点包至最新版本。)
|
||||
</details>
|
||||
|
||||
**v1.0.7**
|
||||
<details>
|
||||
<summary><b>v1.0.7</b></summary>
|
||||
|
||||
- 增加 `easy cascadeLoader` - stable cascade 加载器
|
||||
- 增加 `easy preSamplingCascade` - stabled cascade stage_c 预采样参数
|
||||
- 增加 `easy fullCascadeKSampler` - stable cascade stage_c 完整版采样器
|
||||
- 增加 `easy cascadeKSampler` - stable cascade stage-c ksampler simple
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>v1.0.6</b></summary>
|
||||
@@ -260,38 +347,10 @@
|
||||
| easy poseEditor | [ComfyUI_Custom_Nodes_AlekPet](https://github.com/AlekPet/ComfyUI_Custom_Nodes_AlekPet) | poseNode |
|
||||
| easy if | [ComfyUI-Logic](https://github.com/theUpsider/ComfyUI-Logic) | IfExecute |
|
||||
| easy preSamplingLayerDiffusion | [ComfyUI-layerdiffusion](https://github.com/huchenlei/ComfyUI-layerdiffusion) | LayeredDiffusionApply等 |
|
||||
|
||||
## 示例
|
||||
|
||||
导入后请自行更换您目录里的大模型
|
||||
|
||||
### StableDiffusion
|
||||
#### 文生图
|
||||
|
||||
<img src="https://raw.githubusercontent.com/yolain/yolain-comfyui-workflow/main/Workflows/Simple/text_to_image.png">
|
||||
|
||||
#### 图生图+controlnet
|
||||
|
||||
<img src="https://raw.githubusercontent.com/yolain/yolain-comfyui-workflow/main/Workflows/Simple/image_to_image_controlnet.png">
|
||||
|
||||
#### InstantID
|
||||
|
||||
<img src="https://raw.githubusercontent.com/yolain/yolain-comfyui-workflow/main/Workflows/Simple/instantID.png">
|
||||
|
||||
### LayerDiffusion
|
||||
#### SD15
|
||||
<img src="https://raw.githubusercontent.com/yolain/yolain-comfyui-workflow/main/Workflows/Simple/layer_diffusion_sd15.png">
|
||||
|
||||
#### SDXL
|
||||
<img src="https://raw.githubusercontent.com/yolain/yolain-comfyui-workflow/main/Workflows/Simple/layer_diffusion_example.png">
|
||||
|
||||
### StableCascade
|
||||
#### 文生图
|
||||
<img src="https://raw.githubusercontent.com/yolain/yolain-comfyui-workflow/main/Workflows/StableCascade/text_to_image.png">
|
||||
|
||||
#### 图生图
|
||||
<img src="https://raw.githubusercontent.com/yolain/yolain-comfyui-workflow/main/Workflows/StableCascade/image_to_image.png">
|
||||
|
||||
| easy dynamiCrafterLoader | [ComfyUI-layerdiffusion](https://github.com/ExponentialML/ComfyUI_Native_DynamiCrafter) | Apply Dynamicrafter |
|
||||
| easy imageChooser | [cg-image-picker](https://github.com/chrisgoringe/cg-image-picker) | Preview Chooser |
|
||||
| easy styleAlignedBatchAlign | [style_aligned_comfy](https://github.com/chrisgoringe/cg-image-picker) | styleAlignedBatchAlign |
|
||||
| easy icLightApply | [ComfyUI-IC-Light](https://github.com/huchenlei/ComfyUI-IC-Light) | ICLightApply等 |
|
||||
|
||||
## Credits
|
||||
|
||||
@@ -308,3 +367,15 @@
|
||||
[ComfyUI-Inspire-Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) - 常规整合包2
|
||||
|
||||
[ComfyUI-Logic](https://github.com/theUpsider/ComfyUI-Logic) - ComfyUI逻辑运算
|
||||
|
||||
[ComfyUI-ResAdapter](https://github.com/jiaxiangc/ComfyUI-ResAdapter) - 让模型生成不受训练分辨率限制
|
||||
|
||||
[ComfyUI_IPAdapter_plus](https://github.com/cubiq/ComfyUI_IPAdapter_plus) - 风格迁移
|
||||
|
||||
[ComfyUI_InstantID](https://github.com/cubiq/ComfyUI_InstantID) - 人脸迁移
|
||||
|
||||
[ComfyUI-Custom-Scripts](https://github.com/pythongosssss/ComfyUI-Custom-Scripts) - pyssss 小蛇🐍脚本
|
||||
|
||||
[cg-image-picker](https://github.com/chrisgoringe/cg-image-picker) - 图片选择器
|
||||
|
||||
[ComfyUI-BrushNet](https://github.com/nullquant/ComfyUI-BrushNet) - BrushNet 内补节点
|
||||
|
||||
+15
-31
@@ -1,7 +1,9 @@
|
||||
__version__ = "1.1.8"
|
||||
|
||||
import os
|
||||
import glob
|
||||
import folder_paths
|
||||
import importlib
|
||||
from pathlib import Path
|
||||
|
||||
node_list = [
|
||||
"server",
|
||||
@@ -23,7 +25,7 @@ cwd_path = os.path.dirname(os.path.realpath(__file__))
|
||||
comfy_path = folder_paths.base_path
|
||||
|
||||
#Wildcards读取
|
||||
from .py.wildcards import read_wildcard_dict
|
||||
from .py.libs.wildcards import read_wildcard_dict
|
||||
wildcards_path = os.path.join(os.path.dirname(__file__), "wildcards")
|
||||
if os.path.exists(wildcards_path):
|
||||
read_wildcard_dict(wildcards_path)
|
||||
@@ -40,37 +42,19 @@ else:
|
||||
os.mkdir(styles_path)
|
||||
os.mkdir(samples_path)
|
||||
|
||||
#合并autocomplete覆盖到pyssss包
|
||||
pyssss_path = os.path.join(comfy_path, "custom_nodes", "ComfyUI-Custom-Scripts", "user")
|
||||
combine_folder = os.path.join(cwd_path, "autocomplete")
|
||||
if os.path.exists(combine_folder):
|
||||
pass
|
||||
else:
|
||||
os.mkdir(combine_folder)
|
||||
if os.path.exists(pyssss_path):
|
||||
output_file = os.path.join(pyssss_path, "autocomplete.txt")
|
||||
# 遍历 combine 目录下的所有 txt 文件,读取内容并合并
|
||||
merged_content = ''
|
||||
for file_path in glob.glob(os.path.join(combine_folder, '*.txt')):
|
||||
with open(file_path, 'r', encoding='utf-8', errors='ignore') as file:
|
||||
try:
|
||||
file_content = file.read()
|
||||
merged_content += file_content + '\n'
|
||||
except UnicodeDecodeError:
|
||||
pass
|
||||
# 备份之前的autocomplete
|
||||
# bak_file = os.path.join(pyssss_path, "autocomplete.txt.bak")
|
||||
# if os.path.exists(bak_file):
|
||||
# pass
|
||||
# elif os.path.exists(output_file):
|
||||
# shutil.copy(output_file, bak_file)
|
||||
if merged_content != '':
|
||||
# 将合并的内容写入目标文件 autocomplete.txt,并指定编码为 utf-8
|
||||
with open(output_file, 'w', encoding='utf-8') as target_file:
|
||||
target_file.write(merged_content)
|
||||
# ComfyUI-Easy-PS相关 (需要把模型预览图暴露给PS读取,此处借鉴了 AIGODLIKE-ComfyUI-Studio 的部分代码)
|
||||
from .py.libs.add_resources import add_static_resource
|
||||
from .py.libs.model import easyModelManager
|
||||
model_config = easyModelManager().models_config
|
||||
for model in model_config:
|
||||
paths = folder_paths.get_folder_paths(model)
|
||||
for path in paths:
|
||||
if not Path(path).exists():
|
||||
continue
|
||||
add_static_resource(path, path, limit=True)
|
||||
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', "WEB_DIRECTORY"]
|
||||
|
||||
|
||||
print('\033[34mComfy-Easy-Use (v1.1.1): \033[92mLoaded\033[0m')
|
||||
print(f'\033[34mComfy-Easy-Use v{__version__}: \033[92mLoaded\033[0m')
|
||||
+16
@@ -0,0 +1,16 @@
|
||||
@echo off
|
||||
|
||||
set "requirements_txt=%~dp0\requirements.txt"
|
||||
set "python_exec=..\..\..\python_embeded\python.exe"
|
||||
|
||||
echo Installing EasyUse Requirements...
|
||||
|
||||
if exist "%python_exec%" (
|
||||
echo Installing with ComfyUI Portable
|
||||
"%python_exec%" -s -m pip install -r "%requirements_txt%"
|
||||
) else (
|
||||
echo Installing with system Python
|
||||
pip install -r "%requirements_txt%"
|
||||
)
|
||||
|
||||
pause
|
||||
@@ -0,0 +1,34 @@
|
||||
import folder_paths
|
||||
import os
|
||||
def add_folder_path_and_extensions(folder_name, full_folder_paths, extensions):
|
||||
for full_folder_path in full_folder_paths:
|
||||
folder_paths.add_model_folder_path(folder_name, full_folder_path)
|
||||
if folder_name in folder_paths.folder_names_and_paths:
|
||||
current_paths, current_extensions = folder_paths.folder_names_and_paths[folder_name]
|
||||
updated_extensions = current_extensions | extensions
|
||||
folder_paths.folder_names_and_paths[folder_name] = (current_paths, updated_extensions)
|
||||
else:
|
||||
folder_paths.folder_names_and_paths[folder_name] = (full_folder_paths, extensions)
|
||||
|
||||
image_suffixs = set([".jpg", ".jpeg", ".png", ".gif", ".webp", ".bmp", ".tiff", ".svg", ".ico", ".apng", ".tif", ".hdr", ".exr"])
|
||||
|
||||
model_path = folder_paths.models_dir
|
||||
add_folder_path_and_extensions("ultralytics_bbox", [os.path.join(model_path, "ultralytics", "bbox")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("ultralytics_segm", [os.path.join(model_path, "ultralytics", "segm")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("ultralytics", [os.path.join(model_path, "ultralytics")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("mmdets_bbox", [os.path.join(model_path, "mmdets", "bbox")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("mmdets_segm", [os.path.join(model_path, "mmdets", "segm")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("mmdets", [os.path.join(model_path, "mmdets")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("sams", [os.path.join(model_path, "sams")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("onnx", [os.path.join(model_path, "onnx")], {'.onnx'})
|
||||
add_folder_path_and_extensions("instantid", [os.path.join(model_path, "instantid")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("layer_model", [os.path.join(model_path, "layer_model")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("rembg", [os.path.join(model_path, "rembg")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("ipadapter", [os.path.join(model_path, "ipadapter")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("dynamicrafter_models", [os.path.join(model_path, "dynamicrafter_models")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("mediapipe", [os.path.join(model_path, "mediapipe")], set(['.tflite','.pth']))
|
||||
add_folder_path_and_extensions("inpaint", [os.path.join(model_path, "inpaint")], folder_paths.supported_pt_extensions)
|
||||
add_folder_path_and_extensions("prompt_generator", [os.path.join(model_path, "prompt_generator")], folder_paths.supported_pt_extensions)
|
||||
|
||||
add_folder_path_and_extensions("checkpoints_thumb", [os.path.join(model_path, "checkpoints")], image_suffixs)
|
||||
add_folder_path_and_extensions("loras_thumb", [os.path.join(model_path, "loras")], image_suffixs)
|
||||
@@ -1,12 +1,17 @@
|
||||
import re
|
||||
import os
|
||||
import torch
|
||||
import hashlib
|
||||
import sys
|
||||
import json
|
||||
import shutil
|
||||
import folder_paths
|
||||
from folder_paths import get_directory_by_type
|
||||
from server import PromptServer
|
||||
from .config import RESOURCES_DIR, FOOOCUS_STYLES_DIR, FOOOCUS_STYLES_SAMPLES
|
||||
from .easyNodes import easyCache
|
||||
from .logic import ConvertAnything
|
||||
from .libs.model import easyModelManager
|
||||
from .libs.utils import getMetadata, cleanGPUUsedForce, get_local_filepath
|
||||
from .libs.cache import remove_cache
|
||||
from .libs.translate import has_chinese, zh_to_en
|
||||
|
||||
try:
|
||||
import aiohttp
|
||||
@@ -16,6 +21,34 @@ except ImportError:
|
||||
print("pip install aiohttp")
|
||||
sys.exit()
|
||||
|
||||
@PromptServer.instance.routes.post("/easyuse/cleangpu")
|
||||
def cleanGPU(request):
|
||||
try:
|
||||
cleanGPUUsedForce()
|
||||
remove_cache('*')
|
||||
return web.Response(status=200)
|
||||
except Exception as e:
|
||||
return web.Response(status=500)
|
||||
pass
|
||||
|
||||
@PromptServer.instance.routes.post("/easyuse/translate")
|
||||
async def translate(request):
|
||||
post = await request.post()
|
||||
text = post.get("text")
|
||||
if has_chinese(text):
|
||||
return web.json_response({"text": zh_to_en([text])[0]})
|
||||
else:
|
||||
return web.json_response({"text": text})
|
||||
|
||||
@PromptServer.instance.routes.get("/easyuse/reboot")
|
||||
def reboot(request):
|
||||
try:
|
||||
sys.stdout.close_log()
|
||||
except Exception as e:
|
||||
pass
|
||||
|
||||
return os.execv(sys.executable, [sys.executable] + sys.argv)
|
||||
|
||||
# parse csv
|
||||
@PromptServer.instance.routes.post("/easyuse/upload/csv")
|
||||
async def parse_csv(request):
|
||||
@@ -86,6 +119,7 @@ async def getStylesImage(request):
|
||||
return web.Response(text=FOOOCUS_STYLES_SAMPLES + name + '.jpg')
|
||||
return web.Response(status=400)
|
||||
|
||||
# convert type
|
||||
@PromptServer.instance.routes.post("/easyuse/convert")
|
||||
async def convertType(request):
|
||||
post = await request.post()
|
||||
@@ -97,7 +131,168 @@ async def convertType(request):
|
||||
else:
|
||||
return web.Response(status=400)
|
||||
|
||||
# get models lists
|
||||
@PromptServer.instance.routes.get("/easyuse/models/list")
|
||||
async def getModelsList(request):
|
||||
if "type" in request.rel_url.query:
|
||||
type = request.rel_url.query["type"]
|
||||
if type not in ['checkpoints', 'loras']:
|
||||
return web.Response(status=400)
|
||||
manager = easyModelManager()
|
||||
return web.json_response(manager.get_model_lists(type))
|
||||
else:
|
||||
return web.Response(status=400)
|
||||
|
||||
# get models thumbnails
|
||||
@PromptServer.instance.routes.get("/easyuse/models/thumbnail")
|
||||
async def getModelsThumbnail(request):
|
||||
checkpoints = folder_paths.get_filename_list("checkpoints_thumb")
|
||||
loras = folder_paths.get_filename_list("loras_thumb")
|
||||
checkpoints_full = []
|
||||
loras_full = []
|
||||
if len(checkpoints) + len(loras) >= 500:
|
||||
return web.Response(status=400)
|
||||
for index, i in enumerate(checkpoints):
|
||||
full_path = folder_paths.get_full_path('checkpoints_thumb', str(i))
|
||||
if full_path:
|
||||
checkpoints_full.append(full_path)
|
||||
for index, i in enumerate(loras):
|
||||
full_path = folder_paths.get_full_path('loras_thumb', str(i))
|
||||
if full_path:
|
||||
loras_full.append(full_path)
|
||||
return web.json_response(checkpoints_full + loras_full)
|
||||
|
||||
@PromptServer.instance.routes.post("/easyuse/metadata/notes/{name}")
|
||||
async def save_notes(request):
|
||||
name = request.match_info["name"]
|
||||
pos = name.index("/")
|
||||
type = name[0:pos]
|
||||
name = name[pos+1:]
|
||||
|
||||
file_path = None
|
||||
if type == "embeddings" or type == "loras":
|
||||
name = name.lower()
|
||||
files = folder_paths.get_filename_list(type)
|
||||
for f in files:
|
||||
lower_f = f.lower()
|
||||
if lower_f == name:
|
||||
file_path = folder_paths.get_full_path(type, f)
|
||||
else:
|
||||
n = os.path.splitext(f)[0].lower()
|
||||
if n == name:
|
||||
file_path = folder_paths.get_full_path(type, f)
|
||||
|
||||
if file_path is not None:
|
||||
break
|
||||
else:
|
||||
file_path = folder_paths.get_full_path(
|
||||
type, name)
|
||||
if not file_path:
|
||||
return web.Response(status=404)
|
||||
|
||||
file_no_ext = os.path.splitext(file_path)[0]
|
||||
info_file = file_no_ext + ".txt"
|
||||
with open(info_file, "w") as f:
|
||||
f.write(await request.text())
|
||||
|
||||
return web.Response(status=200)
|
||||
|
||||
@PromptServer.instance.routes.get("/easyuse/metadata/{name}")
|
||||
async def load_metadata(request):
|
||||
name = request.match_info["name"]
|
||||
pos = name.index("/")
|
||||
type = name[0:pos]
|
||||
name = name[pos+1:]
|
||||
|
||||
file_path = None
|
||||
if type == "embeddings":
|
||||
name = name.lower()
|
||||
files = folder_paths.get_filename_list(type)
|
||||
for f in files:
|
||||
lower_f = f.lower()
|
||||
if lower_f == name:
|
||||
file_path = folder_paths.get_full_path(type, f)
|
||||
else:
|
||||
n = os.path.splitext(f)[0].lower()
|
||||
if n == name:
|
||||
file_path = folder_paths.get_full_path(type, f)
|
||||
|
||||
if file_path is not None:
|
||||
break
|
||||
else:
|
||||
file_path = folder_paths.get_full_path(type, name)
|
||||
if not file_path:
|
||||
return web.Response(status=404)
|
||||
|
||||
try:
|
||||
header = getMetadata(file_path)
|
||||
header_json = json.loads(header)
|
||||
meta = header_json["__metadata__"] if "__metadata__" in header_json else None
|
||||
except:
|
||||
meta = None
|
||||
|
||||
if meta is None:
|
||||
meta = {}
|
||||
|
||||
file_no_ext = os.path.splitext(file_path)[0]
|
||||
|
||||
info_file = file_no_ext + ".txt"
|
||||
if os.path.isfile(info_file):
|
||||
with open(info_file, "r") as f:
|
||||
meta["easyuse.notes"] = f.read()
|
||||
|
||||
hash_file = file_no_ext + ".sha256"
|
||||
if os.path.isfile(hash_file):
|
||||
with open(hash_file, "rt") as f:
|
||||
meta["easyuse.sha256"] = f.read()
|
||||
else:
|
||||
with open(file_path, "rb") as f:
|
||||
meta["easyuse.sha256"] = hashlib.sha256(f.read()).hexdigest()
|
||||
with open(hash_file, "wt") as f:
|
||||
f.write(meta["easyuse.sha256"])
|
||||
|
||||
return web.json_response(meta)
|
||||
|
||||
@PromptServer.instance.routes.post("/easyuse/save/{name}")
|
||||
async def save_preview(request):
|
||||
name = request.match_info["name"]
|
||||
pos = name.index("/")
|
||||
type = name[0:pos]
|
||||
name = name[pos+1:]
|
||||
|
||||
body = await request.json()
|
||||
|
||||
dir = get_directory_by_type(body.get("type", "output"))
|
||||
subfolder = body.get("subfolder", "")
|
||||
full_output_folder = os.path.join(dir, os.path.normpath(subfolder))
|
||||
|
||||
if os.path.commonpath((dir, os.path.abspath(full_output_folder))) != dir:
|
||||
return web.Response(status=400)
|
||||
|
||||
filepath = os.path.join(full_output_folder, body.get("filename", ""))
|
||||
image_path = folder_paths.get_full_path(type, name)
|
||||
image_path = os.path.splitext(
|
||||
image_path)[0] + os.path.splitext(filepath)[1]
|
||||
|
||||
shutil.copyfile(filepath, image_path)
|
||||
|
||||
return web.json_response({
|
||||
"image": type + "/" + os.path.basename(image_path)
|
||||
})
|
||||
|
||||
@PromptServer.instance.routes.post("/easyuse/model/download")
|
||||
async def download_model(request):
|
||||
post = await request.post()
|
||||
url = post.get("url")
|
||||
local_dir = post.get("local_dir")
|
||||
if local_dir not in ['checkpoints', 'loras', 'controlnet', 'onnx', 'instantid', 'ipadapter', 'dynamicrafter_models', 'mediapipe', 'rembg', 'layer_model']:
|
||||
return web.Response(status=400)
|
||||
local_path = os.path.join(folder_paths.models_dir, local_dir)
|
||||
try:
|
||||
get_local_filepath(url, local_path)
|
||||
return web.Response(status=200)
|
||||
except:
|
||||
return web.Response(status=500)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
+187
-9
@@ -38,7 +38,7 @@ MAX_SEED_NUM = 1125899906842624
|
||||
|
||||
RESOURCES_DIR = os.path.join(Path(__file__).parent.parent, "resources")
|
||||
|
||||
# fooocus
|
||||
# inpaint
|
||||
INPAINT_DIR = os.path.join(folder_paths.models_dir, "inpaint")
|
||||
FOOOCUS_STYLES_DIR = os.path.join(Path(__file__).parent.parent, "styles")
|
||||
FOOOCUS_STYLES_SAMPLES = 'https://raw.githubusercontent.com/lllyasviel/Fooocus/main/sdxl_styles/samples/'
|
||||
@@ -58,6 +58,29 @@ FOOOCUS_INPAINT_PATCH = {
|
||||
"model_url": "https://huggingface.co/lllyasviel/fooocus_inpaint/resolve/main/inpaint.fooocus.patch"
|
||||
},
|
||||
}
|
||||
BRUSHNET_MODELS = {
|
||||
"random_mask": {
|
||||
"sd1": {
|
||||
"model_url": "https://huggingface.co/Kijai/BrushNet-fp16/resolve/main/brushnet_random_mask_fp16.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": "https://huggingface.co/yolain/brushnet/resolve/main/brushnet_random_mask_sdxl.safetensors"
|
||||
}
|
||||
},
|
||||
"segmentation_mask": {
|
||||
"sd1": {
|
||||
"model_url": "https://huggingface.co/Kijai/BrushNet-fp16/resolve/main/brushnet_segmentation_mask_fp16.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": "https://huggingface.co/yolain/brushnet/resolve/main/brushnet_segmentation_mask_sdxl.safetensors"
|
||||
}
|
||||
}
|
||||
}
|
||||
POWERPAINT_CLIP = {
|
||||
"base_fp16":{
|
||||
"model_url":"https://huggingface.co/runwayml/stable-diffusion-v1-5/resolve/main/text_encoder/model.fp16.safetensors"
|
||||
}
|
||||
}
|
||||
|
||||
# layerDiffuse
|
||||
LAYER_DIFFUSION_DIR = os.path.join(folder_paths.models_dir, "layer_model")
|
||||
@@ -68,7 +91,7 @@ LAYER_DIFFUSION_VAE = {
|
||||
}
|
||||
},
|
||||
"decode": {
|
||||
"sd15": {
|
||||
"sd1": {
|
||||
"model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_vae_transparent_decoder.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
@@ -78,7 +101,7 @@ LAYER_DIFFUSION_VAE = {
|
||||
}
|
||||
LAYER_DIFFUSION = {
|
||||
"Attention Injection": {
|
||||
"sd15": {
|
||||
"sd1": {
|
||||
"model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_transparent_attn.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
@@ -89,12 +112,12 @@ LAYER_DIFFUSION = {
|
||||
"sdxl": {
|
||||
"model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_xl_transparent_conv.safetensors"
|
||||
},
|
||||
"sd15": {
|
||||
"sd1": {
|
||||
"model_url": None
|
||||
}
|
||||
},
|
||||
"Everything": {
|
||||
"sd15": {
|
||||
"sd1": {
|
||||
"model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_joint.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
@@ -102,7 +125,7 @@ LAYER_DIFFUSION = {
|
||||
}
|
||||
},
|
||||
"Foreground": {
|
||||
"sd15": {
|
||||
"sd1": {
|
||||
"model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_fg2bg.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
@@ -110,7 +133,7 @@ LAYER_DIFFUSION = {
|
||||
}
|
||||
},
|
||||
"Foreground to Background": {
|
||||
"sd15": {
|
||||
"sd1": {
|
||||
"model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_fg2bg.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
@@ -118,7 +141,7 @@ LAYER_DIFFUSION = {
|
||||
}
|
||||
},
|
||||
"Background": {
|
||||
"sd15": {
|
||||
"sd1": {
|
||||
"model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_bg2fg.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
@@ -126,7 +149,7 @@ LAYER_DIFFUSION = {
|
||||
}
|
||||
},
|
||||
"Background to Foreground": {
|
||||
"sd15": {
|
||||
"sd1": {
|
||||
"model_url": "https://huggingface.co/LayerDiffusion/layerdiffusion-v1/resolve/main/layer_sd15_bg2fg.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
@@ -135,10 +158,165 @@ LAYER_DIFFUSION = {
|
||||
},
|
||||
}
|
||||
|
||||
# IC Light
|
||||
IC_LIGHT_MODELS = {
|
||||
"Foreground": {
|
||||
"sd1": {
|
||||
"model_url": "https://huggingface.co/huchenlei/IC-Light-ldm/resolve/main/iclight_sd15_fc_unet_ldm.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": None
|
||||
}
|
||||
},
|
||||
"Foreground&Background": {
|
||||
"sd1": {
|
||||
"model_url": "https://huggingface.co/huchenlei/IC-Light-ldm/resolve/main/iclight_sd15_fbc_unet_ldm.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": None
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
# REMBG
|
||||
REMBG_DIR = os.path.join(folder_paths.models_dir, "rembg")
|
||||
REMBG_MODELS = {
|
||||
"RMBG-1.4": {
|
||||
"model_url": "https://huggingface.co/briaai/RMBG-1.4/resolve/main/model.pth"
|
||||
}
|
||||
}
|
||||
|
||||
#ipadapter
|
||||
IPADAPTER_DIR = os.path.join(folder_paths.models_dir, "ipadapter")
|
||||
IPADAPTER_MODELS = {
|
||||
"LIGHT - SD1.5 only (low strength)": {
|
||||
"sd15": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter_sd15_light_v11.bin"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": ""
|
||||
}
|
||||
},
|
||||
"STANDARD (medium strength)": {
|
||||
"sd15": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter_sd15.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/sdxl_models/ip-adapter_sdxl_vit-h.safetensors"
|
||||
}
|
||||
},
|
||||
"VIT-G (medium strength)": {
|
||||
"sd15": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter_sd15_vit-G.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/sdxl_models/ip-adapter_sdxl.safetensors"
|
||||
}
|
||||
},
|
||||
"PLUS (high strength)": {
|
||||
"sd15": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter-plus_sd15.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/sdxl_models/ip-adapter-plus_sdxl_vit-h.safetensors"
|
||||
}
|
||||
},
|
||||
"PLUS FACE (portraits)": {
|
||||
"sd15": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter-plus-face_sd15.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/sdxl_models/ip-adapter-plus-face_sdxl_vit-h.safetensors"
|
||||
}
|
||||
},
|
||||
"FULL FACE - SD1.5 only (portraits stronger)": {
|
||||
"sd15": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter/resolve/main/models/ip-adapter-full-face_sd15.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": ""
|
||||
}
|
||||
},
|
||||
"FACEID": {
|
||||
"sd15": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid_sd15.bin",
|
||||
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid_sd15_lora.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid_sdxl.bin",
|
||||
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid_sdxl_lora.safetensors"
|
||||
}
|
||||
},
|
||||
"FACEID PLUS - SD1.5 only": {
|
||||
"sd15": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plus_sd15.bin",
|
||||
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plus_sd15_lora.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": "",
|
||||
"lora_url": ""
|
||||
}
|
||||
},
|
||||
"FACEID PLUS V2": {
|
||||
"sd15": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plusv2_sd15.bin",
|
||||
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plusv2_sd15_lora.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plusv2_sdxl.bin",
|
||||
"lora_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plusv2_sdxl_lora.safetensors"
|
||||
}
|
||||
},
|
||||
"FACEID PORTRAIT (style transfer)": {
|
||||
"sd15": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-portrait-v11_sd15.bin",
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": "https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-portrait_sdxl.bin",
|
||||
}
|
||||
},
|
||||
"COMPOSITION": {
|
||||
"sd15": {
|
||||
"model_url": "https://huggingface.co/ostris/ip-composition-adapter/resolve/main/ip_plus_composition_sd15.safetensors"
|
||||
},
|
||||
"sdxl": {
|
||||
"model_url": "https://huggingface.co/ostris/ip-composition-adapter/resolve/main/ip_plus_composition_sdxl.safetensors"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
# dynamiCrafter
|
||||
DYNAMICRAFTER_DIR = os.path.join(folder_paths.models_dir, "dynamicrafter_models")
|
||||
DYNAMICRAFTER_MODELS = {
|
||||
"dynamicrafter_unet_512 (2.98GB)": {
|
||||
"model_url": "https://huggingface.co/ExponentialML/DynamiCrafterUNet/resolve/main/dynamicrafter_unet_512.safetensors",
|
||||
"vae_url": "https://huggingface.co/stabilityai/sd-vae-ft-mse-original/resolve/main/vae-ft-mse-840000-ema-pruned.safetensors",
|
||||
"clip_url": "https://huggingface.co/stabilityai/stable-diffusion-2-1/resolve/main/text_encoder/model.safetensors",
|
||||
"clip_vision_url": "https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/resolve/main/open_clip_pytorch_model.safetensors",
|
||||
},
|
||||
"dynamicrafter_unet_512_interp (2.98GB)": {
|
||||
"model_url": "https://huggingface.co/ExponentialML/DynamiCrafterUNet/resolve/main/dynamicrafter_unet_512_interp.safetensors"
|
||||
},
|
||||
"dynamicrafter_unet_1024 (2.98GB)": {
|
||||
"model_url": "https://huggingface.co/ExponentialML/DynamiCrafterUNet/resolve/main/dynamicrafter_unet_1024.safetensors"
|
||||
},
|
||||
"dynamicrafter_unet_256 (2.98GB)": {
|
||||
"model_url": "https://huggingface.co/ExponentialML/DynamiCrafterUNet/resolve/main/dynamicrafter_unet_256.safetensors"
|
||||
},
|
||||
}
|
||||
|
||||
#humanParsing
|
||||
HUMANPARSING_MODELS = {
|
||||
"parsing_lip": {
|
||||
"model_url": "https://huggingface.co/levihsu/OOTDiffusion/resolve/main/checkpoints/humanparsing/parsing_lip.onnx",
|
||||
},
|
||||
}
|
||||
|
||||
#mediapipe
|
||||
MEDIAPIPE_DIR = os.path.join(folder_paths.models_dir, "mediapipe")
|
||||
MEDIAPIPE_MODELS = {
|
||||
"selfie_multiclass_256x256": {
|
||||
"model_url": "https://huggingface.co/yolain/selfie_multiclass_256x256/resolve/main/selfie_multiclass_256x256.tflite"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,332 @@
|
||||
import os
|
||||
import torch
|
||||
import comfy
|
||||
|
||||
from einops import rearrange
|
||||
from comfy import model_base, model_management
|
||||
from .lvdm.modules.networks.openaimodel3d import UNetModel as DynamiCrafterUNetModel
|
||||
|
||||
from .utils.model_utils import DynamiCrafterBase, DYNAMICRAFTER_CONFIG, load_image_proj_dict, load_dynamicrafter_dict, get_image_proj_model
|
||||
|
||||
class DynamiCrafter:
|
||||
|
||||
def __init__(self):
|
||||
self.model_patcher = None
|
||||
|
||||
# There is probably a better way to do this, but with the apply_model callback, this seems necessary.
|
||||
# The model gets wrapped around a CFG Denoiser class, and handles the conditioning parts there.
|
||||
# We cannot access it, so we must find the conditioning according to how ComfyUI handles it.
|
||||
def get_conditioning_pair(self, c_crossattn, use_cfg: bool):
|
||||
if not use_cfg:
|
||||
return c_crossattn
|
||||
|
||||
conditioning_group = []
|
||||
|
||||
for i in range(c_crossattn.shape[0]):
|
||||
# Get the positive and negative conditioning.
|
||||
positive_idx = i + 1
|
||||
negative_idx = i
|
||||
|
||||
if positive_idx >= c_crossattn.shape[0]:
|
||||
break
|
||||
|
||||
if not torch.equal(c_crossattn[[positive_idx]], c_crossattn[[negative_idx]]):
|
||||
conditioning_group = [
|
||||
c_crossattn[[positive_idx]],
|
||||
c_crossattn[[negative_idx]]
|
||||
]
|
||||
break
|
||||
|
||||
if len(conditioning_group) == 0:
|
||||
raise ValueError("Could not get the appropriate conditioning group.")
|
||||
|
||||
return torch.cat(conditioning_group)
|
||||
|
||||
# apply_model, {"input": input_x, "timestep": timestep_, "c": c, "cond_or_uncond": cond_or_uncond}
|
||||
def _forward(self, *args):
|
||||
transformer_options = self.model_patcher.model_options['transformer_options']
|
||||
conditioning = transformer_options['conditioning']
|
||||
|
||||
apply_model = args[0]
|
||||
|
||||
# forward_dict
|
||||
fd = args[1]
|
||||
|
||||
x, t, model_in_kwargs, _ = fd['input'], fd['timestep'], fd['c'], fd['cond_or_uncond']
|
||||
|
||||
c_crossattn = model_in_kwargs.pop("c_crossattn")
|
||||
c_concat = conditioning['c_concat']
|
||||
num_video_frames = conditioning['num_video_frames']
|
||||
fs = conditioning['fs']
|
||||
|
||||
original_num_frames = num_video_frames
|
||||
|
||||
# Better way to determine if we're using CFG
|
||||
# The cond batch will always be num_frames >= 2 since we're doing video,
|
||||
# so we need get this condition differently here.
|
||||
if x.shape[0] > num_video_frames:
|
||||
num_video_frames *= 2
|
||||
batch_size = 2
|
||||
use_cfg = True
|
||||
else:
|
||||
use_cfg = False
|
||||
batch_size = 1
|
||||
|
||||
if use_cfg:
|
||||
c_concat = torch.cat([c_concat] * 2)
|
||||
|
||||
self.validate_forwardable_latent(x, c_concat, num_video_frames, use_cfg)
|
||||
|
||||
x_in, c_concat = map(lambda xc: rearrange(xc, '(b t) c h w -> b c t h w', b=batch_size), (x, c_concat))
|
||||
|
||||
# We always assume video, so there will always be batched conditionings.
|
||||
c_crossattn = self.get_conditioning_pair(c_crossattn, use_cfg)
|
||||
c_crossattn = c_crossattn[:2] if use_cfg else c_crossattn[:1]
|
||||
context_in = c_crossattn
|
||||
|
||||
img_embs = conditioning['image_emb']
|
||||
|
||||
if use_cfg:
|
||||
img_emb_uncond = conditioning['image_emb_uncond']
|
||||
img_embs = torch.cat([img_embs, img_emb_uncond])
|
||||
|
||||
fs = torch.cat([fs] * x_in.shape[0])
|
||||
|
||||
outs = []
|
||||
for i in range(batch_size):
|
||||
model_in_kwargs['transformer_options']['cond_idx'] = i
|
||||
x_out = apply_model(
|
||||
x_in[[i]],
|
||||
t=torch.cat([t[:1]]),
|
||||
context_in=context_in[[i]],
|
||||
c_crossattn=c_crossattn,
|
||||
cc_concat=c_concat[[i]], # "cc" is to handle naming conflict with apply_model wrapper.
|
||||
# We want to handle this in the UNet forward.
|
||||
num_video_frames=num_video_frames // 2 if batch_size > 1 else num_video_frames,
|
||||
img_emb=img_embs[[i]],
|
||||
fs=fs[[i]],
|
||||
**model_in_kwargs
|
||||
)
|
||||
outs.append(x_out)
|
||||
|
||||
x_out = torch.cat(list(reversed(outs)))
|
||||
x_out = rearrange(x_out, 'b c t h w -> (b t) c h w')
|
||||
|
||||
return x_out
|
||||
|
||||
def assign_forward_args(
|
||||
self,
|
||||
model,
|
||||
c_concat,
|
||||
image_emb,
|
||||
image_emb_uncond,
|
||||
fs,
|
||||
frames,
|
||||
):
|
||||
model.model_options['transformer_options']['conditioning'] = {
|
||||
"c_concat": c_concat,
|
||||
"image_emb": image_emb,
|
||||
'image_emb_uncond': image_emb_uncond,
|
||||
"fs": fs,
|
||||
"num_video_frames": frames,
|
||||
}
|
||||
|
||||
def validate_forwardable_latent(self, latent, c_concat, num_video_frames, use_cfg):
|
||||
check_no_cfg = latent.shape[0] != num_video_frames
|
||||
check_with_cfg = latent.shape[0] != (num_video_frames * 2)
|
||||
|
||||
latent_batch_size = latent.shape[0] if not use_cfg else latent.shape[0] // 2
|
||||
num_frames = num_video_frames if not use_cfg else num_video_frames // 2
|
||||
|
||||
if all([check_no_cfg, check_with_cfg]):
|
||||
raise ValueError(
|
||||
"Please make sure your latent inputs match the number of frames in the DynamiCrafter Processor."
|
||||
f"Got a latent batch size of ({latent_batch_size}) with number of frames being ({num_frames})."
|
||||
)
|
||||
|
||||
latent_h, latent_w = latent.shape[-2:]
|
||||
c_concat_h, c_concat_w = c_concat.shape[-2:]
|
||||
|
||||
if not all([latent_h == c_concat_h, latent_w == c_concat_w]):
|
||||
raise ValueError(
|
||||
"Please make sure that your input latent and image frames are the same height and width.",
|
||||
f"Image Size: {c_concat_w * 8}, {c_concat_h * 8}, Latent Size: {latent_h * 8}, {latent_w * 8}"
|
||||
)
|
||||
|
||||
def process_image_conditioning(
|
||||
self,
|
||||
model,
|
||||
clip_vision,
|
||||
vae,
|
||||
image_proj_model,
|
||||
images,
|
||||
use_interpolate,
|
||||
fps: int,
|
||||
frames: int,
|
||||
scale_latents: bool
|
||||
):
|
||||
self.model_patcher = model
|
||||
encoded_latent = vae.encode(images[:, :, :, :3])
|
||||
|
||||
encoded_image = clip_vision.encode_image(images[:1])['last_hidden_state']
|
||||
image_emb = image_proj_model(encoded_image)
|
||||
|
||||
encoded_image_uncond = clip_vision.encode_image(torch.zeros_like(images)[:1])['last_hidden_state']
|
||||
image_emb_uncond = image_proj_model(encoded_image_uncond)
|
||||
|
||||
c_concat = encoded_latent
|
||||
|
||||
if scale_latents:
|
||||
vae_process_input = vae.process_input
|
||||
vae.process_input = lambda image: (image - .5) * 2
|
||||
c_concat = vae.encode(images[:, :, :, :3])
|
||||
vae.process_input = vae_process_input
|
||||
c_concat = model.model.process_latent_in(c_concat) * 1.3
|
||||
else:
|
||||
c_concat = model.model.process_latent_in(c_concat)
|
||||
|
||||
fs = torch.tensor([fps], dtype=torch.long, device=model_management.intermediate_device())
|
||||
|
||||
model.set_model_unet_function_wrapper(self._forward)
|
||||
|
||||
used_interpolate_processing = False
|
||||
|
||||
if use_interpolate and frames > 16:
|
||||
raise ValueError(
|
||||
"When using interpolation mode, the maximum amount of frames are 16."
|
||||
"If you're doing long video generation, consider using the last frame\
|
||||
from the first generation for the next one (autoregressive)."
|
||||
)
|
||||
if encoded_latent.shape[0] == 1:
|
||||
c_concat = torch.cat([c_concat] * frames, dim=0)[:frames]
|
||||
|
||||
if use_interpolate:
|
||||
mask = torch.zeros_like(c_concat)
|
||||
mask[:1] = c_concat[:1]
|
||||
c_concat = mask
|
||||
|
||||
used_interpolate_processing = True
|
||||
else:
|
||||
if use_interpolate and c_concat.shape[0] in [2, 3]:
|
||||
input_frame_count = c_concat.shape[0]
|
||||
|
||||
# We're just padding to the same type an size of the concat
|
||||
masked_frames = torch.zeros_like(torch.cat([c_concat[:1]] * frames))[:frames]
|
||||
|
||||
# Start frame
|
||||
masked_frames[:1] = c_concat[:1]
|
||||
|
||||
end_frame_idx = -1
|
||||
|
||||
# TODO
|
||||
speed = 1.0
|
||||
if speed < 1.0:
|
||||
possible_speeds = list(torch.linspace(0, 1.0, c_concat.shape[0]))
|
||||
speed_from_frames = enumerate(possible_speeds)
|
||||
speed_idx = min(speed_from_frames, key=lambda n: n[1] - speed)[0]
|
||||
end_frame_idx = speed_idx
|
||||
|
||||
# End frame
|
||||
masked_frames[-1:] = c_concat[[end_frame_idx]]
|
||||
|
||||
# Possible middle frame, but not working at the moment.
|
||||
if input_frame_count == 3:
|
||||
middle_idx = masked_frames.shape[0] // 2
|
||||
middle_idx_frame = c_concat.shape[0] // 2
|
||||
masked_frames[[middle_idx]] = c_concat[[middle_idx_frame]]
|
||||
|
||||
c_concat = masked_frames
|
||||
used_interpolate_processing = True
|
||||
|
||||
print(f"Using interpolation mode with {input_frame_count} frames.")
|
||||
|
||||
if c_concat.shape[0] < frames and not used_interpolate_processing:
|
||||
print(
|
||||
"Multiple images found, but interpolation mode is unset. Using the first frame as condition.",
|
||||
)
|
||||
c_concat = torch.cat([c_concat[:1]] * frames)
|
||||
|
||||
c_concat = c_concat[:frames]
|
||||
|
||||
if encoded_latent.shape[0] == 1:
|
||||
encoded_latent = torch.cat([encoded_latent] * frames)[:frames]
|
||||
|
||||
if encoded_latent.shape[0] < frames and encoded_latent.shape[0] != 1:
|
||||
encoded_latent = torch.cat(
|
||||
[encoded_latent] + [encoded_latent[-1:]] * abs(encoded_latent.shape[0] - frames)
|
||||
)[:frames]
|
||||
|
||||
# We could store this as a state in this Node Class Instance, but to prevent any weird edge cases,
|
||||
# this should always be passed through the 'stateless' way, and let ComfyUI handle the transformer_options state.
|
||||
self.assign_forward_args(model, c_concat, image_emb, image_emb_uncond, fs, frames)
|
||||
|
||||
return (model, {"samples": torch.zeros_like(c_concat)}, {"samples": encoded_latent},)
|
||||
|
||||
|
||||
# Loader for the DynamiCrafter model.
|
||||
def load_model_sicts(self, model_path: str):
|
||||
model_state_dict = comfy.utils.load_torch_file(model_path)
|
||||
dynamicrafter_dict = load_dynamicrafter_dict(model_state_dict)
|
||||
image_proj_dict = load_image_proj_dict(model_state_dict)
|
||||
|
||||
return dynamicrafter_dict, image_proj_dict
|
||||
|
||||
def get_prediction_type(self, is_eps: bool, model_config):
|
||||
if not is_eps and "image_cross_attention_scale_learnable" in model_config.unet_config.keys():
|
||||
model_config.unet_config["image_cross_attention_scale_learnable"] = False
|
||||
|
||||
return model_base.ModelType.EPS if is_eps else model_base.ModelType.V_PREDICTION
|
||||
|
||||
def handle_model_management(self, dynamicrafter_dict: dict, model_config):
|
||||
parameters = comfy.utils.calculate_parameters(dynamicrafter_dict, "model.diffusion_model.")
|
||||
load_device = model_management.get_torch_device()
|
||||
unet_dtype = model_management.unet_dtype(
|
||||
model_params=parameters,
|
||||
supported_dtypes=model_config.supported_inference_dtypes
|
||||
)
|
||||
manual_cast_dtype = model_management.unet_manual_cast(
|
||||
unet_dtype,
|
||||
load_device,
|
||||
model_config.supported_inference_dtypes
|
||||
)
|
||||
model_config.set_inference_dtype(unet_dtype, manual_cast_dtype)
|
||||
inital_load_device = model_management.unet_inital_load_device(parameters, unet_dtype)
|
||||
offload_device = model_management.unet_offload_device()
|
||||
|
||||
return load_device, inital_load_device
|
||||
|
||||
def check_leftover_keys(self, state_dict: dict):
|
||||
left_over = state_dict.keys()
|
||||
if len(left_over) > 0:
|
||||
print("left over keys:", left_over)
|
||||
|
||||
def load_dynamicrafter(self, model_path):
|
||||
|
||||
if os.path.exists(model_path):
|
||||
dynamicrafter_dict, image_proj_dict = self.load_model_sicts(model_path)
|
||||
model_config = DynamiCrafterBase(DYNAMICRAFTER_CONFIG)
|
||||
|
||||
dynamicrafter_dict, is_eps = model_config.process_dict_version(state_dict=dynamicrafter_dict)
|
||||
|
||||
MODEL_TYPE = self.get_prediction_type(is_eps, model_config)
|
||||
load_device, inital_load_device = self.handle_model_management(dynamicrafter_dict, model_config)
|
||||
|
||||
model = model_base.BaseModel(
|
||||
model_config,
|
||||
model_type=MODEL_TYPE,
|
||||
device=inital_load_device,
|
||||
unet_model=DynamiCrafterUNetModel
|
||||
)
|
||||
|
||||
image_proj_model = get_image_proj_model(image_proj_dict)
|
||||
model.load_model_weights(dynamicrafter_dict, "model.diffusion_model.")
|
||||
self.check_leftover_keys(dynamicrafter_dict)
|
||||
|
||||
model_patcher = comfy.model_patcher.ModelPatcher(
|
||||
model,
|
||||
load_device=load_device,
|
||||
offload_device=model_management.unet_offload_device(),
|
||||
current_device=inital_load_device
|
||||
)
|
||||
|
||||
return (model_patcher, image_proj_model,)
|
||||
@@ -0,0 +1,102 @@
|
||||
# adopted from
|
||||
# https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py
|
||||
# and
|
||||
# https://github.com/lucidrains/denoising-diffusion-pytorch/blob/7706bdfc6f527f58d33f84b7b522e61e6e3164b3/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py
|
||||
# and
|
||||
# https://github.com/openai/guided-diffusion/blob/0ba878e517b276c45d1195eb29f6f5f72659a05b/guided_diffusion/nn.py
|
||||
#
|
||||
# thanks!
|
||||
|
||||
import torch.nn as nn
|
||||
import comfy.ops
|
||||
ops = comfy.ops.disable_weight_init
|
||||
|
||||
from ..utils.utils import instantiate_from_config
|
||||
|
||||
def disabled_train(self, mode=True):
|
||||
"""Overwrite model.train with this function to make sure train/eval mode
|
||||
does not change anymore."""
|
||||
return self
|
||||
|
||||
def zero_module(module):
|
||||
"""
|
||||
Zero out the parameters of a module and return it.
|
||||
"""
|
||||
for p in module.parameters():
|
||||
p.detach().zero_()
|
||||
return module
|
||||
|
||||
def scale_module(module, scale):
|
||||
"""
|
||||
Scale the parameters of a module and return it.
|
||||
"""
|
||||
for p in module.parameters():
|
||||
p.detach().mul_(scale)
|
||||
return module
|
||||
|
||||
|
||||
def conv_nd(dims, *args, **kwargs):
|
||||
"""
|
||||
Create a 1D, 2D, or 3D convolution module.
|
||||
"""
|
||||
if dims == 1:
|
||||
return nn.Conv1d(*args, **kwargs)
|
||||
elif dims == 2:
|
||||
return ops.Conv2d(*args, **kwargs)
|
||||
elif dims == 3:
|
||||
return ops.Conv3d(*args, **kwargs)
|
||||
raise ValueError(f"unsupported dimensions: {dims}")
|
||||
|
||||
|
||||
def linear(*args, **kwargs):
|
||||
"""
|
||||
Create a linear module.
|
||||
"""
|
||||
return ops.Linear(*args, **kwargs)
|
||||
|
||||
|
||||
def avg_pool_nd(dims, *args, **kwargs):
|
||||
"""
|
||||
Create a 1D, 2D, or 3D average pooling module.
|
||||
"""
|
||||
if dims == 1:
|
||||
return nn.AvgPool1d(*args, **kwargs)
|
||||
elif dims == 2:
|
||||
return nn.AvgPool2d(*args, **kwargs)
|
||||
elif dims == 3:
|
||||
return nn.AvgPool3d(*args, **kwargs)
|
||||
raise ValueError(f"unsupported dimensions: {dims}")
|
||||
|
||||
|
||||
def nonlinearity(type='silu'):
|
||||
if type == 'silu':
|
||||
return nn.SiLU()
|
||||
elif type == 'leaky_relu':
|
||||
return nn.LeakyReLU()
|
||||
|
||||
|
||||
class GroupNormSpecific(ops.GroupNorm):
|
||||
def forward(self, x):
|
||||
return super().forward(x.float()).type(x.dtype)
|
||||
|
||||
|
||||
def normalization(channels, num_groups=32, dtype=None, device=None):
|
||||
"""
|
||||
Make a standard normalization layer.
|
||||
:param channels: number of input channels.
|
||||
:return: an nn.Module for normalization.
|
||||
"""
|
||||
return GroupNormSpecific(num_groups, channels, dtype=dtype, device=device)
|
||||
|
||||
|
||||
class HybridConditioner(nn.Module):
|
||||
|
||||
def __init__(self, c_concat_config, c_crossattn_config):
|
||||
super().__init__()
|
||||
self.concat_conditioner = instantiate_from_config(c_concat_config)
|
||||
self.crossattn_conditioner = instantiate_from_config(c_crossattn_config)
|
||||
|
||||
def forward(self, c_concat, c_crossattn):
|
||||
c_concat = self.concat_conditioner(c_concat)
|
||||
c_crossattn = self.crossattn_conditioner(c_crossattn)
|
||||
return {'c_concat': [c_concat], 'c_crossattn': [c_crossattn]}
|
||||
@@ -0,0 +1,94 @@
|
||||
import math
|
||||
from inspect import isfunction
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
def gather_data(data, return_np=True):
|
||||
''' gather data from multiple processes to one list '''
|
||||
data_list = [torch.zeros_like(data) for _ in range(dist.get_world_size())]
|
||||
dist.all_gather(data_list, data) # gather not supported with NCCL
|
||||
if return_np:
|
||||
data_list = [data.cpu().numpy() for data in data_list]
|
||||
return data_list
|
||||
|
||||
def autocast(f):
|
||||
def do_autocast(*args, **kwargs):
|
||||
with torch.cuda.amp.autocast(enabled=True,
|
||||
dtype=torch.get_autocast_gpu_dtype(),
|
||||
cache_enabled=torch.is_autocast_cache_enabled()):
|
||||
return f(*args, **kwargs)
|
||||
return do_autocast
|
||||
|
||||
|
||||
def extract_into_tensor(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
out = a.gather(-1, t)
|
||||
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
||||
|
||||
|
||||
def noise_like(shape, device, repeat=False):
|
||||
repeat_noise = lambda: torch.randn((1, *shape[1:]), device=device).repeat(shape[0], *((1,) * (len(shape) - 1)))
|
||||
noise = lambda: torch.randn(shape, device=device)
|
||||
return repeat_noise() if repeat else noise()
|
||||
|
||||
|
||||
def default(val, d):
|
||||
if exists(val):
|
||||
return val
|
||||
return d() if isfunction(d) else d
|
||||
|
||||
def exists(val):
|
||||
return val is not None
|
||||
|
||||
def identity(*args, **kwargs):
|
||||
return nn.Identity()
|
||||
|
||||
def uniq(arr):
|
||||
return{el: True for el in arr}.keys()
|
||||
|
||||
def mean_flat(tensor):
|
||||
"""
|
||||
Take the mean over all non-batch dimensions.
|
||||
"""
|
||||
return tensor.mean(dim=list(range(1, len(tensor.shape))))
|
||||
|
||||
def ismap(x):
|
||||
if not isinstance(x, torch.Tensor):
|
||||
return False
|
||||
return (len(x.shape) == 4) and (x.shape[1] > 3)
|
||||
|
||||
def isimage(x):
|
||||
if not isinstance(x,torch.Tensor):
|
||||
return False
|
||||
return (len(x.shape) == 4) and (x.shape[1] == 3 or x.shape[1] == 1)
|
||||
|
||||
def max_neg_value(t):
|
||||
return -torch.finfo(t.dtype).max
|
||||
|
||||
def shape_to_str(x):
|
||||
shape_str = "x".join([str(x) for x in x.shape])
|
||||
return shape_str
|
||||
|
||||
def init_(tensor):
|
||||
dim = tensor.shape[-1]
|
||||
std = 1 / math.sqrt(dim)
|
||||
tensor.uniform_(-std, std)
|
||||
return tensor
|
||||
|
||||
ckpt = torch.utils.checkpoint.checkpoint
|
||||
def checkpoint(func, inputs, params, flag):
|
||||
"""
|
||||
Evaluate a function without caching intermediate activations, allowing for
|
||||
reduced memory at the expense of extra compute in the backward pass.
|
||||
:param func: the function to evaluate.
|
||||
:param inputs: the argument sequence to pass to `func`.
|
||||
:param params: a sequence of parameters `func` depends on but does not
|
||||
explicitly take as arguments.
|
||||
:param flag: if False, disable gradient checkpointing.
|
||||
"""
|
||||
if flag:
|
||||
return ckpt(func, *inputs, use_reentrant=False)
|
||||
else:
|
||||
return func(*inputs)
|
||||
@@ -0,0 +1,95 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
|
||||
class AbstractDistribution:
|
||||
def sample(self):
|
||||
raise NotImplementedError()
|
||||
|
||||
def mode(self):
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class DiracDistribution(AbstractDistribution):
|
||||
def __init__(self, value):
|
||||
self.value = value
|
||||
|
||||
def sample(self):
|
||||
return self.value
|
||||
|
||||
def mode(self):
|
||||
return self.value
|
||||
|
||||
|
||||
class DiagonalGaussianDistribution(object):
|
||||
def __init__(self, parameters, deterministic=False):
|
||||
self.parameters = parameters
|
||||
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
|
||||
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
||||
self.deterministic = deterministic
|
||||
self.std = torch.exp(0.5 * self.logvar)
|
||||
self.var = torch.exp(self.logvar)
|
||||
if self.deterministic:
|
||||
self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)
|
||||
|
||||
def sample(self, noise=None):
|
||||
if noise is None:
|
||||
noise = torch.randn(self.mean.shape)
|
||||
|
||||
x = self.mean + self.std * noise.to(device=self.parameters.device)
|
||||
return x
|
||||
|
||||
def kl(self, other=None):
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.])
|
||||
else:
|
||||
if other is None:
|
||||
return 0.5 * torch.sum(torch.pow(self.mean, 2)
|
||||
+ self.var - 1.0 - self.logvar,
|
||||
dim=[1, 2, 3])
|
||||
else:
|
||||
return 0.5 * torch.sum(
|
||||
torch.pow(self.mean - other.mean, 2) / other.var
|
||||
+ self.var / other.var - 1.0 - self.logvar + other.logvar,
|
||||
dim=[1, 2, 3])
|
||||
|
||||
def nll(self, sample, dims=[1,2,3]):
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.])
|
||||
logtwopi = np.log(2.0 * np.pi)
|
||||
return 0.5 * torch.sum(
|
||||
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
|
||||
dim=dims)
|
||||
|
||||
def mode(self):
|
||||
return self.mean
|
||||
|
||||
|
||||
def normal_kl(mean1, logvar1, mean2, logvar2):
|
||||
"""
|
||||
source: https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/losses.py#L12
|
||||
Compute the KL divergence between two gaussians.
|
||||
Shapes are automatically broadcasted, so batches can be compared to
|
||||
scalars, among other use cases.
|
||||
"""
|
||||
tensor = None
|
||||
for obj in (mean1, logvar1, mean2, logvar2):
|
||||
if isinstance(obj, torch.Tensor):
|
||||
tensor = obj
|
||||
break
|
||||
assert tensor is not None, "at least one argument must be a Tensor"
|
||||
|
||||
# Force variances to be Tensors. Broadcasting helps convert scalars to
|
||||
# Tensors, but it does not work for torch.exp().
|
||||
logvar1, logvar2 = [
|
||||
x if isinstance(x, torch.Tensor) else torch.tensor(x).to(tensor)
|
||||
for x in (logvar1, logvar2)
|
||||
]
|
||||
|
||||
return 0.5 * (
|
||||
-1.0
|
||||
+ logvar2
|
||||
- logvar1
|
||||
+ torch.exp(logvar1 - logvar2)
|
||||
+ ((mean1 - mean2) ** 2) * torch.exp(-logvar2)
|
||||
)
|
||||
@@ -0,0 +1,76 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
class LitEma(nn.Module):
|
||||
def __init__(self, model, decay=0.9999, use_num_upates=True):
|
||||
super().__init__()
|
||||
if decay < 0.0 or decay > 1.0:
|
||||
raise ValueError('Decay must be between 0 and 1')
|
||||
|
||||
self.m_name2s_name = {}
|
||||
self.register_buffer('decay', torch.tensor(decay, dtype=torch.float32))
|
||||
self.register_buffer('num_updates', torch.tensor(0,dtype=torch.int) if use_num_upates
|
||||
else torch.tensor(-1,dtype=torch.int))
|
||||
|
||||
for name, p in model.named_parameters():
|
||||
if p.requires_grad:
|
||||
#remove as '.'-character is not allowed in buffers
|
||||
s_name = name.replace('.','')
|
||||
self.m_name2s_name.update({name:s_name})
|
||||
self.register_buffer(s_name,p.clone().detach().data)
|
||||
|
||||
self.collected_params = []
|
||||
|
||||
def forward(self,model):
|
||||
decay = self.decay
|
||||
|
||||
if self.num_updates >= 0:
|
||||
self.num_updates += 1
|
||||
decay = min(self.decay,(1 + self.num_updates) / (10 + self.num_updates))
|
||||
|
||||
one_minus_decay = 1.0 - decay
|
||||
|
||||
with torch.no_grad():
|
||||
m_param = dict(model.named_parameters())
|
||||
shadow_params = dict(self.named_buffers())
|
||||
|
||||
for key in m_param:
|
||||
if m_param[key].requires_grad:
|
||||
sname = self.m_name2s_name[key]
|
||||
shadow_params[sname] = shadow_params[sname].type_as(m_param[key])
|
||||
shadow_params[sname].sub_(one_minus_decay * (shadow_params[sname] - m_param[key]))
|
||||
else:
|
||||
assert not key in self.m_name2s_name
|
||||
|
||||
def copy_to(self, model):
|
||||
m_param = dict(model.named_parameters())
|
||||
shadow_params = dict(self.named_buffers())
|
||||
for key in m_param:
|
||||
if m_param[key].requires_grad:
|
||||
m_param[key].data.copy_(shadow_params[self.m_name2s_name[key]].data)
|
||||
else:
|
||||
assert not key in self.m_name2s_name
|
||||
|
||||
def store(self, parameters):
|
||||
"""
|
||||
Save the current parameters for restoring later.
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
temporarily stored.
|
||||
"""
|
||||
self.collected_params = [param.clone() for param in parameters]
|
||||
|
||||
def restore(self, parameters):
|
||||
"""
|
||||
Restore the parameters stored with the `store` method.
|
||||
Useful to validate the model with EMA parameters without affecting the
|
||||
original optimization process. Store the parameters before the
|
||||
`copy_to` method. After validation (or model saving), use this to
|
||||
restore the former parameters.
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
updated with the stored parameters.
|
||||
"""
|
||||
for c_param, param in zip(self.collected_params, parameters):
|
||||
param.data.copy_(c_param.data)
|
||||
@@ -0,0 +1,219 @@
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
import torch
|
||||
import numpy as np
|
||||
from einops import rearrange
|
||||
import torch.nn.functional as F
|
||||
import pytorch_lightning as pl
|
||||
from ...modules.networks.ae_modules import Encoder, Decoder
|
||||
from ...distributions import DiagonalGaussianDistribution
|
||||
from utils.utils import instantiate_from_config
|
||||
|
||||
|
||||
class AutoencoderKL(pl.LightningModule):
|
||||
def __init__(self,
|
||||
ddconfig,
|
||||
lossconfig,
|
||||
embed_dim,
|
||||
ckpt_path=None,
|
||||
ignore_keys=[],
|
||||
image_key="image",
|
||||
colorize_nlabels=None,
|
||||
monitor=None,
|
||||
test=False,
|
||||
logdir=None,
|
||||
input_dim=4,
|
||||
test_args=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.image_key = image_key
|
||||
self.encoder = Encoder(**ddconfig)
|
||||
self.decoder = Decoder(**ddconfig)
|
||||
self.loss = instantiate_from_config(lossconfig)
|
||||
assert ddconfig["double_z"]
|
||||
self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1)
|
||||
self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
|
||||
self.embed_dim = embed_dim
|
||||
self.input_dim = input_dim
|
||||
self.test = test
|
||||
self.test_args = test_args
|
||||
self.logdir = logdir
|
||||
if colorize_nlabels is not None:
|
||||
assert type(colorize_nlabels)==int
|
||||
self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1))
|
||||
if monitor is not None:
|
||||
self.monitor = monitor
|
||||
if ckpt_path is not None:
|
||||
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)
|
||||
if self.test:
|
||||
self.init_test()
|
||||
|
||||
def init_test(self,):
|
||||
self.test = True
|
||||
save_dir = os.path.join(self.logdir, "test")
|
||||
if 'ckpt' in self.test_args:
|
||||
ckpt_name = os.path.basename(self.test_args.ckpt).split('.ckpt')[0] + f'_epoch{self._cur_epoch}'
|
||||
self.root = os.path.join(save_dir, ckpt_name)
|
||||
else:
|
||||
self.root = save_dir
|
||||
if 'test_subdir' in self.test_args:
|
||||
self.root = os.path.join(save_dir, self.test_args.test_subdir)
|
||||
|
||||
self.root_zs = os.path.join(self.root, "zs")
|
||||
self.root_dec = os.path.join(self.root, "reconstructions")
|
||||
self.root_inputs = os.path.join(self.root, "inputs")
|
||||
os.makedirs(self.root, exist_ok=True)
|
||||
|
||||
if self.test_args.save_z:
|
||||
os.makedirs(self.root_zs, exist_ok=True)
|
||||
if self.test_args.save_reconstruction:
|
||||
os.makedirs(self.root_dec, exist_ok=True)
|
||||
if self.test_args.save_input:
|
||||
os.makedirs(self.root_inputs, exist_ok=True)
|
||||
assert(self.test_args is not None)
|
||||
self.test_maximum = getattr(self.test_args, 'test_maximum', None)
|
||||
self.count = 0
|
||||
self.eval_metrics = {}
|
||||
self.decodes = []
|
||||
self.save_decode_samples = 2048
|
||||
|
||||
def init_from_ckpt(self, path, ignore_keys=list()):
|
||||
sd = torch.load(path, map_location="cpu")
|
||||
try:
|
||||
self._cur_epoch = sd['epoch']
|
||||
sd = sd["state_dict"]
|
||||
except:
|
||||
self._cur_epoch = 'null'
|
||||
keys = list(sd.keys())
|
||||
for k in keys:
|
||||
for ik in ignore_keys:
|
||||
if k.startswith(ik):
|
||||
print("Deleting key {} from state_dict.".format(k))
|
||||
del sd[k]
|
||||
self.load_state_dict(sd, strict=False)
|
||||
# self.load_state_dict(sd, strict=True)
|
||||
print(f"Restored from {path}")
|
||||
|
||||
def encode(self, x, **kwargs):
|
||||
|
||||
h = self.encoder(x)
|
||||
moments = self.quant_conv(h)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
return posterior
|
||||
|
||||
def decode(self, z, **kwargs):
|
||||
z = self.post_quant_conv(z)
|
||||
dec = self.decoder(z)
|
||||
return dec
|
||||
|
||||
def forward(self, input, sample_posterior=True):
|
||||
posterior = self.encode(input)
|
||||
if sample_posterior:
|
||||
z = posterior.sample()
|
||||
else:
|
||||
z = posterior.mode()
|
||||
dec = self.decode(z)
|
||||
return dec, posterior
|
||||
|
||||
def get_input(self, batch, k):
|
||||
x = batch[k]
|
||||
if x.dim() == 5 and self.input_dim == 4:
|
||||
b,c,t,h,w = x.shape
|
||||
self.b = b
|
||||
self.t = t
|
||||
x = rearrange(x, 'b c t h w -> (b t) c h w')
|
||||
|
||||
return x
|
||||
|
||||
def training_step(self, batch, batch_idx, optimizer_idx):
|
||||
inputs = self.get_input(batch, self.image_key)
|
||||
reconstructions, posterior = self(inputs)
|
||||
|
||||
if optimizer_idx == 0:
|
||||
# train encoder+decoder+logvar
|
||||
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
|
||||
last_layer=self.get_last_layer(), split="train")
|
||||
self.log("aeloss", aeloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
|
||||
self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=False)
|
||||
return aeloss
|
||||
|
||||
if optimizer_idx == 1:
|
||||
# train the discriminator
|
||||
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step,
|
||||
last_layer=self.get_last_layer(), split="train")
|
||||
|
||||
self.log("discloss", discloss, prog_bar=True, logger=True, on_step=True, on_epoch=True)
|
||||
self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=False)
|
||||
return discloss
|
||||
|
||||
def validation_step(self, batch, batch_idx):
|
||||
inputs = self.get_input(batch, self.image_key)
|
||||
reconstructions, posterior = self(inputs)
|
||||
aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, 0, self.global_step,
|
||||
last_layer=self.get_last_layer(), split="val")
|
||||
|
||||
discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, 1, self.global_step,
|
||||
last_layer=self.get_last_layer(), split="val")
|
||||
|
||||
self.log("val/rec_loss", log_dict_ae["val/rec_loss"])
|
||||
self.log_dict(log_dict_ae)
|
||||
self.log_dict(log_dict_disc)
|
||||
return self.log_dict
|
||||
|
||||
def configure_optimizers(self):
|
||||
lr = self.learning_rate
|
||||
opt_ae = torch.optim.Adam(list(self.encoder.parameters())+
|
||||
list(self.decoder.parameters())+
|
||||
list(self.quant_conv.parameters())+
|
||||
list(self.post_quant_conv.parameters()),
|
||||
lr=lr, betas=(0.5, 0.9))
|
||||
opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(),
|
||||
lr=lr, betas=(0.5, 0.9))
|
||||
return [opt_ae, opt_disc], []
|
||||
|
||||
def get_last_layer(self):
|
||||
return self.decoder.conv_out.weight
|
||||
|
||||
@torch.no_grad()
|
||||
def log_images(self, batch, only_inputs=False, **kwargs):
|
||||
log = dict()
|
||||
x = self.get_input(batch, self.image_key)
|
||||
x = x.to(self.device)
|
||||
if not only_inputs:
|
||||
xrec, posterior = self(x)
|
||||
if x.shape[1] > 3:
|
||||
# colorize with random projection
|
||||
assert xrec.shape[1] > 3
|
||||
x = self.to_rgb(x)
|
||||
xrec = self.to_rgb(xrec)
|
||||
log["samples"] = self.decode(torch.randn_like(posterior.sample()))
|
||||
log["reconstructions"] = xrec
|
||||
log["inputs"] = x
|
||||
return log
|
||||
|
||||
def to_rgb(self, x):
|
||||
assert self.image_key == "segmentation"
|
||||
if not hasattr(self, "colorize"):
|
||||
self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x))
|
||||
x = F.conv2d(x, weight=self.colorize)
|
||||
x = 2.*(x-x.min())/(x.max()-x.min()) - 1.
|
||||
return x
|
||||
|
||||
class IdentityFirstStage(torch.nn.Module):
|
||||
def __init__(self, *args, vq_interface=False, **kwargs):
|
||||
self.vq_interface = vq_interface # TODO: Should be true by default but check to not break older stuff
|
||||
super().__init__()
|
||||
|
||||
def encode(self, x, *args, **kwargs):
|
||||
return x
|
||||
|
||||
def decode(self, x, *args, **kwargs):
|
||||
return x
|
||||
|
||||
def quantize(self, x, *args, **kwargs):
|
||||
if self.vq_interface:
|
||||
return x, None, [None, None, None]
|
||||
return x
|
||||
|
||||
def forward(self, x, *args, **kwargs):
|
||||
return x
|
||||
@@ -0,0 +1,762 @@
|
||||
"""
|
||||
wild mixture of
|
||||
https://github.com/openai/improved-diffusion/blob/e94489283bb876ac1477d5dd7709bbbd2d9902ce/improved_diffusion/gaussian_diffusion.py
|
||||
https://github.com/lucidrains/denoising-diffusion-pytorch/blob/7706bdfc6f527f58d33f84b7b522e61e6e3164b3/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py
|
||||
https://github.com/CompVis/taming-transformers
|
||||
-- merci
|
||||
"""
|
||||
|
||||
from functools import partial
|
||||
from contextlib import contextmanager
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
from einops import rearrange, repeat
|
||||
import logging
|
||||
mainlogger = logging.getLogger('mainlogger')
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torchvision.utils import make_grid
|
||||
|
||||
from ...utils.utils import instantiate_from_config
|
||||
from ..ema import LitEma
|
||||
from ..distributions import DiagonalGaussianDistribution
|
||||
from ..models.utils_diffusion import make_beta_schedule, rescale_zero_terminal_snr
|
||||
from ..basics import disabled_train
|
||||
from ..common import (
|
||||
extract_into_tensor,
|
||||
noise_like,
|
||||
exists,
|
||||
default
|
||||
)
|
||||
|
||||
__conditioning_keys__ = {'concat': 'c_concat',
|
||||
'crossattn': 'c_crossattn',
|
||||
'adm': 'y'}
|
||||
|
||||
class DDPM(nn.Module):
|
||||
# classic DDPM with Gaussian diffusion, in image space
|
||||
def __init__(self,
|
||||
unet_config,
|
||||
timesteps=1000,
|
||||
beta_schedule="linear",
|
||||
loss_type="l2",
|
||||
ckpt_path=None,
|
||||
ignore_keys=[],
|
||||
load_only_unet=False,
|
||||
monitor=None,
|
||||
use_ema=True,
|
||||
first_stage_key="image",
|
||||
image_size=256,
|
||||
channels=3,
|
||||
log_every_t=100,
|
||||
clip_denoised=True,
|
||||
linear_start=1e-4,
|
||||
linear_end=2e-2,
|
||||
cosine_s=8e-3,
|
||||
given_betas=None,
|
||||
original_elbo_weight=0.,
|
||||
v_posterior=0., # weight for choosing posterior variance as sigma = (1-v) * beta_tilde + v * beta
|
||||
l_simple_weight=1.,
|
||||
conditioning_key=None,
|
||||
parameterization="eps", # all assuming fixed variance schedules
|
||||
scheduler_config=None,
|
||||
use_positional_encodings=False,
|
||||
learn_logvar=False,
|
||||
logvar_init=0.,
|
||||
rescale_betas_zero_snr=False,
|
||||
):
|
||||
super().__init__()
|
||||
assert parameterization in ["eps", "x0", "v"], 'currently only supporting "eps" and "x0" and "v"'
|
||||
self.parameterization = parameterization
|
||||
mainlogger.info(f"{self.__class__.__name__}: Running in {self.parameterization}-prediction mode")
|
||||
self.cond_stage_model = None
|
||||
self.clip_denoised = clip_denoised
|
||||
self.log_every_t = log_every_t
|
||||
self.first_stage_key = first_stage_key
|
||||
self.channels = channels
|
||||
self.temporal_length = unet_config.params.temporal_length
|
||||
self.image_size = image_size # try conv?
|
||||
if isinstance(self.image_size, int):
|
||||
self.image_size = [self.image_size, self.image_size]
|
||||
self.use_positional_encodings = use_positional_encodings
|
||||
self.model = DiffusionWrapper(unet_config, conditioning_key)
|
||||
#count_params(self.model, verbose=True)
|
||||
self.use_ema = use_ema
|
||||
self.rescale_betas_zero_snr = rescale_betas_zero_snr
|
||||
if self.use_ema:
|
||||
self.model_ema = LitEma(self.model)
|
||||
mainlogger.info(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.")
|
||||
|
||||
self.use_scheduler = scheduler_config is not None
|
||||
if self.use_scheduler:
|
||||
self.scheduler_config = scheduler_config
|
||||
|
||||
self.v_posterior = v_posterior
|
||||
self.original_elbo_weight = original_elbo_weight
|
||||
self.l_simple_weight = l_simple_weight
|
||||
|
||||
if monitor is not None:
|
||||
self.monitor = monitor
|
||||
if ckpt_path is not None:
|
||||
self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys, only_model=load_only_unet)
|
||||
|
||||
self.register_schedule(given_betas=given_betas, beta_schedule=beta_schedule, timesteps=timesteps,
|
||||
linear_start=linear_start, linear_end=linear_end, cosine_s=cosine_s)
|
||||
|
||||
self.loss_type = loss_type
|
||||
|
||||
self.learn_logvar = learn_logvar
|
||||
self.logvar = torch.full(fill_value=logvar_init, size=(self.num_timesteps,))
|
||||
if self.learn_logvar:
|
||||
self.logvar = nn.Parameter(self.logvar, requires_grad=True)
|
||||
|
||||
def register_schedule(self, given_betas=None, beta_schedule="linear", timesteps=1000,
|
||||
linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3):
|
||||
if exists(given_betas):
|
||||
betas = given_betas
|
||||
else:
|
||||
betas = make_beta_schedule(beta_schedule, timesteps, linear_start=linear_start, linear_end=linear_end,
|
||||
cosine_s=cosine_s)
|
||||
if self.rescale_betas_zero_snr:
|
||||
betas = rescale_zero_terminal_snr(betas)
|
||||
|
||||
alphas = 1. - betas
|
||||
alphas_cumprod = np.cumprod(alphas, axis=0)
|
||||
alphas_cumprod_prev = np.append(1., alphas_cumprod[:-1])
|
||||
|
||||
timesteps, = betas.shape
|
||||
self.num_timesteps = int(timesteps)
|
||||
self.linear_start = linear_start
|
||||
self.linear_end = linear_end
|
||||
assert alphas_cumprod.shape[0] == self.num_timesteps, 'alphas have to be defined for each timestep'
|
||||
|
||||
to_torch = partial(torch.tensor, dtype=torch.float32)
|
||||
|
||||
self.register_buffer('betas', to_torch(betas))
|
||||
self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod))
|
||||
self.register_buffer('alphas_cumprod_prev', to_torch(alphas_cumprod_prev))
|
||||
|
||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
||||
self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod)))
|
||||
self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod)))
|
||||
self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod)))
|
||||
|
||||
if self.parameterization != 'v':
|
||||
self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod)))
|
||||
self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod - 1)))
|
||||
else:
|
||||
self.register_buffer('sqrt_recip_alphas_cumprod', torch.zeros_like(to_torch(alphas_cumprod)))
|
||||
self.register_buffer('sqrt_recipm1_alphas_cumprod', torch.zeros_like(to_torch(alphas_cumprod)))
|
||||
|
||||
# calculations for posterior q(x_{t-1} | x_t, x_0)
|
||||
posterior_variance = (1 - self.v_posterior) * betas * (1. - alphas_cumprod_prev) / (
|
||||
1. - alphas_cumprod) + self.v_posterior * betas
|
||||
# above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t)
|
||||
self.register_buffer('posterior_variance', to_torch(posterior_variance))
|
||||
# below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain
|
||||
self.register_buffer('posterior_log_variance_clipped', to_torch(np.log(np.maximum(posterior_variance, 1e-20))))
|
||||
self.register_buffer('posterior_mean_coef1', to_torch(
|
||||
betas * np.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod)))
|
||||
self.register_buffer('posterior_mean_coef2', to_torch(
|
||||
(1. - alphas_cumprod_prev) * np.sqrt(alphas) / (1. - alphas_cumprod)))
|
||||
|
||||
if self.parameterization == "eps":
|
||||
lvlb_weights = self.betas ** 2 / (
|
||||
2 * self.posterior_variance * to_torch(alphas) * (1 - self.alphas_cumprod))
|
||||
elif self.parameterization == "x0":
|
||||
lvlb_weights = 0.5 * np.sqrt(torch.Tensor(alphas_cumprod)) / (2. * 1 - torch.Tensor(alphas_cumprod))
|
||||
elif self.parameterization == "v":
|
||||
lvlb_weights = torch.ones_like(self.betas ** 2 / (
|
||||
2 * self.posterior_variance * to_torch(alphas) * (1 - self.alphas_cumprod)))
|
||||
else:
|
||||
raise NotImplementedError("mu not supported")
|
||||
# TODO how to choose this term
|
||||
lvlb_weights[0] = lvlb_weights[1]
|
||||
self.register_buffer('lvlb_weights', lvlb_weights, persistent=False)
|
||||
assert not torch.isnan(self.lvlb_weights).all()
|
||||
|
||||
@contextmanager
|
||||
def ema_scope(self, context=None):
|
||||
if self.use_ema:
|
||||
self.model_ema.store(self.model.parameters())
|
||||
self.model_ema.copy_to(self.model)
|
||||
if context is not None:
|
||||
mainlogger.info(f"{context}: Switched to EMA weights")
|
||||
try:
|
||||
yield None
|
||||
finally:
|
||||
if self.use_ema:
|
||||
self.model_ema.restore(self.model.parameters())
|
||||
if context is not None:
|
||||
mainlogger.info(f"{context}: Restored training weights")
|
||||
|
||||
def init_from_ckpt(self, path, ignore_keys=list(), only_model=False):
|
||||
sd = torch.load(path, map_location="cpu")
|
||||
if "state_dict" in list(sd.keys()):
|
||||
sd = sd["state_dict"]
|
||||
keys = list(sd.keys())
|
||||
for k in keys:
|
||||
for ik in ignore_keys:
|
||||
if k.startswith(ik):
|
||||
mainlogger.info("Deleting key {} from state_dict.".format(k))
|
||||
del sd[k]
|
||||
missing, unexpected = self.load_state_dict(sd, strict=False) if not only_model else self.model.load_state_dict(
|
||||
sd, strict=False)
|
||||
mainlogger.info(f"Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys")
|
||||
if len(missing) > 0:
|
||||
mainlogger.info(f"Missing Keys: {missing}")
|
||||
if len(unexpected) > 0:
|
||||
mainlogger.info(f"Unexpected Keys: {unexpected}")
|
||||
|
||||
def q_mean_variance(self, x_start, t):
|
||||
"""
|
||||
Get the distribution q(x_t | x_0).
|
||||
:param x_start: the [N x C x ...] tensor of noiseless inputs.
|
||||
:param t: the number of diffusion steps (minus 1). Here, 0 means one step.
|
||||
:return: A tuple (mean, variance, log_variance), all of x_start's shape.
|
||||
"""
|
||||
mean = (extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start)
|
||||
variance = extract_into_tensor(1.0 - self.alphas_cumprod, t, x_start.shape)
|
||||
log_variance = extract_into_tensor(self.log_one_minus_alphas_cumprod, t, x_start.shape)
|
||||
return mean, variance, log_variance
|
||||
|
||||
def predict_start_from_noise(self, x_t, t, noise):
|
||||
return (
|
||||
extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t -
|
||||
extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * noise
|
||||
)
|
||||
|
||||
def predict_start_from_z_and_v(self, x_t, t, v):
|
||||
# self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod)))
|
||||
# self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod)))
|
||||
return (
|
||||
extract_into_tensor(self.sqrt_alphas_cumprod, t, x_t.shape) * x_t -
|
||||
extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_t.shape) * v
|
||||
)
|
||||
|
||||
def predict_eps_from_z_and_v(self, x_t, t, v):
|
||||
return (
|
||||
extract_into_tensor(self.sqrt_alphas_cumprod, t, x_t.shape) * v +
|
||||
extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_t.shape) * x_t
|
||||
)
|
||||
|
||||
def q_posterior(self, x_start, x_t, t):
|
||||
posterior_mean = (
|
||||
extract_into_tensor(self.posterior_mean_coef1, t, x_t.shape) * x_start +
|
||||
extract_into_tensor(self.posterior_mean_coef2, t, x_t.shape) * x_t
|
||||
)
|
||||
posterior_variance = extract_into_tensor(self.posterior_variance, t, x_t.shape)
|
||||
posterior_log_variance_clipped = extract_into_tensor(self.posterior_log_variance_clipped, t, x_t.shape)
|
||||
return posterior_mean, posterior_variance, posterior_log_variance_clipped
|
||||
|
||||
def p_mean_variance(self, x, t, clip_denoised: bool):
|
||||
model_out = self.model(x, t)
|
||||
if self.parameterization == "eps":
|
||||
x_recon = self.predict_start_from_noise(x, t=t, noise=model_out)
|
||||
elif self.parameterization == "x0":
|
||||
x_recon = model_out
|
||||
if clip_denoised:
|
||||
x_recon.clamp_(-1., 1.)
|
||||
|
||||
model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start=x_recon, x_t=x, t=t)
|
||||
return model_mean, posterior_variance, posterior_log_variance
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample(self, x, t, clip_denoised=True, repeat_noise=False):
|
||||
b, *_, device = *x.shape, x.device
|
||||
model_mean, _, model_log_variance = self.p_mean_variance(x=x, t=t, clip_denoised=clip_denoised)
|
||||
noise = noise_like(x.shape, device, repeat_noise)
|
||||
# no noise when t == 0
|
||||
nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1)))
|
||||
return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample_loop(self, shape, return_intermediates=False):
|
||||
device = self.betas.device
|
||||
b = shape[0]
|
||||
img = torch.randn(shape, device=device)
|
||||
intermediates = [img]
|
||||
for i in tqdm(reversed(range(0, self.num_timesteps)), desc='Sampling t', total=self.num_timesteps):
|
||||
img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long),
|
||||
clip_denoised=self.clip_denoised)
|
||||
if i % self.log_every_t == 0 or i == self.num_timesteps - 1:
|
||||
intermediates.append(img)
|
||||
if return_intermediates:
|
||||
return img, intermediates
|
||||
return img
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self, batch_size=16, return_intermediates=False):
|
||||
image_size = self.image_size
|
||||
channels = self.channels
|
||||
return self.p_sample_loop((batch_size, channels, image_size, image_size),
|
||||
return_intermediates=return_intermediates)
|
||||
|
||||
def q_sample(self, x_start, t, noise=None):
|
||||
noise = default(noise, lambda: torch.randn_like(x_start))
|
||||
return (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)
|
||||
|
||||
def get_v(self, x, noise, t):
|
||||
return (
|
||||
extract_into_tensor(self.sqrt_alphas_cumprod, t, x.shape) * noise -
|
||||
extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x.shape) * x
|
||||
)
|
||||
|
||||
def get_input(self, batch, k):
|
||||
x = batch[k]
|
||||
x = x.to(memory_format=torch.contiguous_format).float()
|
||||
return x
|
||||
|
||||
def _get_rows_from_list(self, samples):
|
||||
n_imgs_per_row = len(samples)
|
||||
denoise_grid = rearrange(samples, 'n b c h w -> b n c h w')
|
||||
denoise_grid = rearrange(denoise_grid, 'b n c h w -> (b n) c h w')
|
||||
denoise_grid = make_grid(denoise_grid, nrow=n_imgs_per_row)
|
||||
return denoise_grid
|
||||
|
||||
@torch.no_grad()
|
||||
def log_images(self, batch, N=8, n_row=2, sample=True, return_keys=None, **kwargs):
|
||||
log = dict()
|
||||
x = self.get_input(batch, self.first_stage_key)
|
||||
N = min(x.shape[0], N)
|
||||
n_row = min(x.shape[0], n_row)
|
||||
x = x.to(self.device)[:N]
|
||||
log["inputs"] = x
|
||||
|
||||
# get diffusion row
|
||||
diffusion_row = list()
|
||||
x_start = x[:n_row]
|
||||
|
||||
for t in range(self.num_timesteps):
|
||||
if t % self.log_every_t == 0 or t == self.num_timesteps - 1:
|
||||
t = repeat(torch.tensor([t]), '1 -> b', b=n_row)
|
||||
t = t.to(self.device).long()
|
||||
noise = torch.randn_like(x_start)
|
||||
x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise)
|
||||
diffusion_row.append(x_noisy)
|
||||
|
||||
log["diffusion_row"] = self._get_rows_from_list(diffusion_row)
|
||||
|
||||
if sample:
|
||||
# get denoise row
|
||||
with self.ema_scope("Plotting"):
|
||||
samples, denoise_row = self.sample(batch_size=N, return_intermediates=True)
|
||||
|
||||
log["samples"] = samples
|
||||
log["denoise_row"] = self._get_rows_from_list(denoise_row)
|
||||
|
||||
if return_keys:
|
||||
if np.intersect1d(list(log.keys()), return_keys).shape[0] == 0:
|
||||
return log
|
||||
else:
|
||||
return {key: log[key] for key in return_keys}
|
||||
return log
|
||||
|
||||
|
||||
class LatentDiffusion(DDPM):
|
||||
"""main class"""
|
||||
def __init__(self,
|
||||
first_stage_config,
|
||||
cond_stage_config,
|
||||
num_timesteps_cond=None,
|
||||
cond_stage_key="caption",
|
||||
cond_stage_trainable=False,
|
||||
cond_stage_forward=None,
|
||||
conditioning_key=None,
|
||||
uncond_prob=0.2,
|
||||
uncond_type="empty_seq",
|
||||
scale_factor=1.0,
|
||||
scale_by_std=False,
|
||||
encoder_type="2d",
|
||||
only_model=False,
|
||||
noise_strength=0,
|
||||
use_dynamic_rescale=False,
|
||||
base_scale=0.7,
|
||||
turning_step=400,
|
||||
loop_video=False,
|
||||
fps_condition_type='fs',
|
||||
perframe_ae=False,
|
||||
*args, **kwargs):
|
||||
self.num_timesteps_cond = default(num_timesteps_cond, 1)
|
||||
self.scale_by_std = scale_by_std
|
||||
assert self.num_timesteps_cond <= kwargs['timesteps']
|
||||
# for backwards compatibility after implementation of DiffusionWrapper
|
||||
ckpt_path = kwargs.pop("ckpt_path", None)
|
||||
ignore_keys = kwargs.pop("ignore_keys", [])
|
||||
conditioning_key = default(conditioning_key, 'crossattn')
|
||||
super().__init__(conditioning_key=conditioning_key, *args, **kwargs)
|
||||
|
||||
self.cond_stage_trainable = cond_stage_trainable
|
||||
self.cond_stage_key = cond_stage_key
|
||||
self.noise_strength = noise_strength
|
||||
self.use_dynamic_rescale = use_dynamic_rescale
|
||||
self.loop_video = loop_video
|
||||
self.fps_condition_type = fps_condition_type
|
||||
self.perframe_ae = perframe_ae
|
||||
try:
|
||||
self.num_downs = len(first_stage_config.params.ddconfig.ch_mult) - 1
|
||||
except:
|
||||
self.num_downs = 0
|
||||
if not scale_by_std:
|
||||
self.scale_factor = scale_factor
|
||||
else:
|
||||
self.register_buffer('scale_factor', torch.tensor(scale_factor))
|
||||
|
||||
if use_dynamic_rescale:
|
||||
scale_arr1 = np.linspace(1.0, base_scale, turning_step)
|
||||
scale_arr2 = np.full(self.num_timesteps, base_scale)
|
||||
scale_arr = np.concatenate((scale_arr1, scale_arr2))
|
||||
to_torch = partial(torch.tensor, dtype=torch.float32)
|
||||
self.register_buffer('scale_arr', to_torch(scale_arr))
|
||||
|
||||
self.instantiate_first_stage(first_stage_config)
|
||||
self.instantiate_cond_stage(cond_stage_config)
|
||||
self.first_stage_config = first_stage_config
|
||||
self.cond_stage_config = cond_stage_config
|
||||
self.clip_denoised = False
|
||||
|
||||
self.cond_stage_forward = cond_stage_forward
|
||||
self.encoder_type = encoder_type
|
||||
assert(encoder_type in ["2d", "3d"])
|
||||
self.uncond_prob = uncond_prob
|
||||
self.classifier_free_guidance = True if uncond_prob > 0 else False
|
||||
assert(uncond_type in ["zero_embed", "empty_seq"])
|
||||
self.uncond_type = uncond_type
|
||||
|
||||
self.restarted_from_ckpt = False
|
||||
if ckpt_path is not None:
|
||||
self.init_from_ckpt(ckpt_path, ignore_keys, only_model=only_model)
|
||||
self.restarted_from_ckpt = True
|
||||
|
||||
|
||||
def make_cond_schedule(self, ):
|
||||
self.cond_ids = torch.full(size=(self.num_timesteps,), fill_value=self.num_timesteps - 1, dtype=torch.long)
|
||||
ids = torch.round(torch.linspace(0, self.num_timesteps - 1, self.num_timesteps_cond)).long()
|
||||
self.cond_ids[:self.num_timesteps_cond] = ids
|
||||
|
||||
def instantiate_first_stage(self, config):
|
||||
model = instantiate_from_config(config)
|
||||
self.first_stage_model = model.eval()
|
||||
self.first_stage_model.train = disabled_train
|
||||
for param in self.first_stage_model.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def instantiate_cond_stage(self, config):
|
||||
if not self.cond_stage_trainable:
|
||||
model = instantiate_from_config(config)
|
||||
self.cond_stage_model = model.eval()
|
||||
self.cond_stage_model.train = disabled_train
|
||||
for param in self.cond_stage_model.parameters():
|
||||
param.requires_grad = False
|
||||
else:
|
||||
model = instantiate_from_config(config)
|
||||
self.cond_stage_model = model
|
||||
|
||||
def get_learned_conditioning(self, c):
|
||||
if self.cond_stage_forward is None:
|
||||
if hasattr(self.cond_stage_model, 'encode') and callable(self.cond_stage_model.encode):
|
||||
c = self.cond_stage_model.encode(c)
|
||||
if isinstance(c, DiagonalGaussianDistribution):
|
||||
c = c.mode()
|
||||
else:
|
||||
c = self.cond_stage_model(c)
|
||||
else:
|
||||
assert hasattr(self.cond_stage_model, self.cond_stage_forward)
|
||||
c = getattr(self.cond_stage_model, self.cond_stage_forward)(c)
|
||||
return c
|
||||
|
||||
def get_first_stage_encoding(self, encoder_posterior, noise=None):
|
||||
if isinstance(encoder_posterior, DiagonalGaussianDistribution):
|
||||
z = encoder_posterior.sample(noise=noise)
|
||||
elif isinstance(encoder_posterior, torch.Tensor):
|
||||
z = encoder_posterior
|
||||
else:
|
||||
raise NotImplementedError(f"encoder_posterior of type '{type(encoder_posterior)}' not yet implemented")
|
||||
return self.scale_factor * z
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x):
|
||||
if self.encoder_type == "2d" and x.dim() == 5:
|
||||
b, _, t, _, _ = x.shape
|
||||
x = rearrange(x, 'b c t h w -> (b t) c h w')
|
||||
reshape_back = True
|
||||
else:
|
||||
reshape_back = False
|
||||
|
||||
## consume more GPU memory but faster
|
||||
if not self.perframe_ae:
|
||||
encoder_posterior = self.first_stage_model.encode(x)
|
||||
results = self.get_first_stage_encoding(encoder_posterior).detach()
|
||||
else: ## consume less GPU memory but slower
|
||||
results = []
|
||||
for index in range(x.shape[0]):
|
||||
frame_batch = self.first_stage_model.encode(x[index:index+1,:,:,:])
|
||||
frame_result = self.get_first_stage_encoding(frame_batch).detach()
|
||||
results.append(frame_result)
|
||||
results = torch.cat(results, dim=0)
|
||||
|
||||
if reshape_back:
|
||||
results = rearrange(results, '(b t) c h w -> b c t h w', b=b,t=t)
|
||||
|
||||
return results
|
||||
|
||||
def decode_core(self, z, **kwargs):
|
||||
if self.encoder_type == "2d" and z.dim() == 5:
|
||||
b, _, t, _, _ = z.shape
|
||||
z = rearrange(z, 'b c t h w -> (b t) c h w')
|
||||
reshape_back = True
|
||||
else:
|
||||
reshape_back = False
|
||||
|
||||
if not self.perframe_ae:
|
||||
z = 1. / self.scale_factor * z
|
||||
results = self.first_stage_model.decode(z, **kwargs)
|
||||
else:
|
||||
results = []
|
||||
for index in range(z.shape[0]):
|
||||
frame_z = 1. / self.scale_factor * z[index:index+1,:,:,:]
|
||||
frame_result = self.first_stage_model.decode(frame_z, **kwargs)
|
||||
results.append(frame_result)
|
||||
results = torch.cat(results, dim=0)
|
||||
|
||||
if reshape_back:
|
||||
results = rearrange(results, '(b t) c h w -> b c t h w', b=b,t=t)
|
||||
return results
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z, **kwargs):
|
||||
return self.decode_core(z, **kwargs)
|
||||
|
||||
# same as above but without decorator
|
||||
def differentiable_decode_first_stage(self, z, **kwargs):
|
||||
return self.decode_core(z, **kwargs)
|
||||
|
||||
def forward(self, x, c, **kwargs):
|
||||
t = torch.randint(0, self.num_timesteps, (x.shape[0],), device=self.device).long()
|
||||
if self.use_dynamic_rescale:
|
||||
x = x * extract_into_tensor(self.scale_arr, t, x.shape)
|
||||
return self.p_losses(x, c, t, **kwargs)
|
||||
|
||||
def apply_model(self, x_noisy, t, cond, **kwargs):
|
||||
if isinstance(cond, dict):
|
||||
# hybrid case, cond is exptected to be a dict
|
||||
pass
|
||||
else:
|
||||
if not isinstance(cond, list):
|
||||
cond = [cond]
|
||||
key = 'c_concat' if self.model.conditioning_key == 'concat' else 'c_crossattn'
|
||||
cond = {key: cond}
|
||||
|
||||
x_recon = self.model(x_noisy, t, **cond, **kwargs)
|
||||
|
||||
if isinstance(x_recon, tuple):
|
||||
return x_recon[0]
|
||||
else:
|
||||
return x_recon
|
||||
|
||||
def _get_denoise_row_from_list(self, samples, desc=''):
|
||||
denoise_row = []
|
||||
for zd in tqdm(samples, desc=desc):
|
||||
denoise_row.append(self.decode_first_stage(zd.to(self.device)))
|
||||
n_log_timesteps = len(denoise_row)
|
||||
|
||||
denoise_row = torch.stack(denoise_row) # n_log_timesteps, b, C, H, W
|
||||
|
||||
if denoise_row.dim() == 5:
|
||||
denoise_grid = rearrange(denoise_row, 'n b c h w -> b n c h w')
|
||||
denoise_grid = rearrange(denoise_grid, 'b n c h w -> (b n) c h w')
|
||||
denoise_grid = make_grid(denoise_grid, nrow=n_log_timesteps)
|
||||
elif denoise_row.dim() == 6:
|
||||
# video, grid_size=[n_log_timesteps*bs, t]
|
||||
video_length = denoise_row.shape[3]
|
||||
denoise_grid = rearrange(denoise_row, 'n b c t h w -> b n c t h w')
|
||||
denoise_grid = rearrange(denoise_grid, 'b n c t h w -> (b n) c t h w')
|
||||
denoise_grid = rearrange(denoise_grid, 'n c t h w -> (n t) c h w')
|
||||
denoise_grid = make_grid(denoise_grid, nrow=video_length)
|
||||
else:
|
||||
raise ValueError
|
||||
|
||||
return denoise_grid
|
||||
|
||||
|
||||
def p_mean_variance(self, x, c, t, clip_denoised: bool, return_x0=False, score_corrector=None, corrector_kwargs=None, **kwargs):
|
||||
t_in = t
|
||||
model_out = self.apply_model(x, t_in, c, **kwargs)
|
||||
|
||||
if score_corrector is not None:
|
||||
assert self.parameterization == "eps"
|
||||
model_out = score_corrector.modify_score(self, model_out, x, t, c, **corrector_kwargs)
|
||||
|
||||
if self.parameterization == "eps":
|
||||
x_recon = self.predict_start_from_noise(x, t=t, noise=model_out)
|
||||
elif self.parameterization == "x0":
|
||||
x_recon = model_out
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
if clip_denoised:
|
||||
x_recon.clamp_(-1., 1.)
|
||||
|
||||
model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start=x_recon, x_t=x, t=t)
|
||||
|
||||
if return_x0:
|
||||
return model_mean, posterior_variance, posterior_log_variance, x_recon
|
||||
else:
|
||||
return model_mean, posterior_variance, posterior_log_variance
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample(self, x, c, t, clip_denoised=False, repeat_noise=False, return_x0=False, \
|
||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, **kwargs):
|
||||
b, *_, device = *x.shape, x.device
|
||||
outputs = self.p_mean_variance(x=x, c=c, t=t, clip_denoised=clip_denoised, return_x0=return_x0, \
|
||||
score_corrector=score_corrector, corrector_kwargs=corrector_kwargs, **kwargs)
|
||||
if return_x0:
|
||||
model_mean, _, model_log_variance, x0 = outputs
|
||||
else:
|
||||
model_mean, _, model_log_variance = outputs
|
||||
|
||||
noise = noise_like(x.shape, device, repeat_noise) * temperature
|
||||
if noise_dropout > 0.:
|
||||
noise = torch.nn.functional.dropout(noise, p=noise_dropout)
|
||||
# no noise when t == 0
|
||||
nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1)))
|
||||
|
||||
if return_x0:
|
||||
return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise, x0
|
||||
else:
|
||||
return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample_loop(self, cond, shape, return_intermediates=False, x_T=None, verbose=True, callback=None, \
|
||||
timesteps=None, mask=None, x0=None, img_callback=None, start_T=None, log_every_t=None, **kwargs):
|
||||
|
||||
if not log_every_t:
|
||||
log_every_t = self.log_every_t
|
||||
device = self.betas.device
|
||||
b = shape[0]
|
||||
# sample an initial noise
|
||||
if x_T is None:
|
||||
img = torch.randn(shape, device=device)
|
||||
else:
|
||||
img = x_T
|
||||
|
||||
intermediates = [img]
|
||||
if timesteps is None:
|
||||
timesteps = self.num_timesteps
|
||||
if start_T is not None:
|
||||
timesteps = min(timesteps, start_T)
|
||||
|
||||
iterator = tqdm(reversed(range(0, timesteps)), desc='Sampling t', total=timesteps) if verbose else reversed(range(0, timesteps))
|
||||
|
||||
if mask is not None:
|
||||
assert x0 is not None
|
||||
assert x0.shape[2:3] == mask.shape[2:3] # spatial size has to match
|
||||
|
||||
for i in iterator:
|
||||
ts = torch.full((b,), i, device=device, dtype=torch.long)
|
||||
if self.shorten_cond_schedule:
|
||||
assert self.model.conditioning_key != 'hybrid'
|
||||
tc = self.cond_ids[ts].to(cond.device)
|
||||
cond = self.q_sample(x_start=cond, t=tc, noise=torch.randn_like(cond))
|
||||
|
||||
img = self.p_sample(img, cond, ts, clip_denoised=self.clip_denoised, **kwargs)
|
||||
if mask is not None:
|
||||
img_orig = self.q_sample(x0, ts)
|
||||
img = img_orig * mask + (1. - mask) * img
|
||||
|
||||
if i % log_every_t == 0 or i == timesteps - 1:
|
||||
intermediates.append(img)
|
||||
if callback: callback(i)
|
||||
if img_callback: img_callback(img, i)
|
||||
|
||||
if return_intermediates:
|
||||
return img, intermediates
|
||||
return img
|
||||
|
||||
|
||||
class LatentVisualDiffusion(LatentDiffusion):
|
||||
def __init__(self, img_cond_stage_config, image_proj_stage_config, freeze_embedder=True, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self._init_embedder(img_cond_stage_config, freeze_embedder)
|
||||
self.image_proj_model = instantiate_from_config(image_proj_stage_config)
|
||||
|
||||
def _init_embedder(self, config, freeze=True):
|
||||
embedder = instantiate_from_config(config)
|
||||
if freeze:
|
||||
self.embedder = embedder.eval()
|
||||
self.embedder.train = disabled_train
|
||||
for param in self.embedder.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
|
||||
class DiffusionWrapper(nn.Module):
|
||||
def __init__(self, diff_model_config, conditioning_key):
|
||||
super().__init__()
|
||||
self.diffusion_model = instantiate_from_config(diff_model_config)
|
||||
self.conditioning_key = conditioning_key
|
||||
|
||||
def forward(self, x, t, c_concat: list = None, c_crossattn: list = None,
|
||||
c_adm=None, s=None, mask=None, **kwargs):
|
||||
# temporal_context = fps is foNone
|
||||
if self.conditioning_key is None:
|
||||
out = self.diffusion_model(x, t)
|
||||
elif self.conditioning_key == 'concat':
|
||||
xc = torch.cat([x] + c_concat, dim=1)
|
||||
out = self.diffusion_model(xc, t, **kwargs)
|
||||
elif self.conditioning_key == 'crossattn':
|
||||
cc = torch.cat(c_crossattn, 1)
|
||||
out = self.diffusion_model(x, t, context=cc, **kwargs)
|
||||
elif self.conditioning_key == 'hybrid':
|
||||
## it is just right [b,c,t,h,w]: concatenate in channel dim
|
||||
xc = torch.cat([x] + c_concat, dim=1)
|
||||
cc = torch.cat(c_crossattn, 1)
|
||||
out = self.diffusion_model(xc, t, context=cc, **kwargs)
|
||||
elif self.conditioning_key == 'resblockcond':
|
||||
cc = c_crossattn[0]
|
||||
out = self.diffusion_model(x, t, context=cc)
|
||||
elif self.conditioning_key == 'adm':
|
||||
cc = c_crossattn[0]
|
||||
out = self.diffusion_model(x, t, y=cc)
|
||||
elif self.conditioning_key == 'hybrid-adm':
|
||||
assert c_adm is not None
|
||||
xc = torch.cat([x] + c_concat, dim=1)
|
||||
cc = torch.cat(c_crossattn, 1)
|
||||
out = self.diffusion_model(xc, t, context=cc, y=c_adm, **kwargs)
|
||||
elif self.conditioning_key == 'hybrid-time':
|
||||
assert s is not None
|
||||
xc = torch.cat([x] + c_concat, dim=1)
|
||||
cc = torch.cat(c_crossattn, 1)
|
||||
out = self.diffusion_model(xc, t, context=cc, s=s)
|
||||
elif self.conditioning_key == 'concat-time-mask':
|
||||
# assert s is not None
|
||||
xc = torch.cat([x] + c_concat, dim=1)
|
||||
out = self.diffusion_model(xc, t, context=None, s=s, mask=mask)
|
||||
elif self.conditioning_key == 'concat-adm-mask':
|
||||
# assert s is not None
|
||||
if c_concat is not None:
|
||||
xc = torch.cat([x] + c_concat, dim=1)
|
||||
else:
|
||||
xc = x
|
||||
out = self.diffusion_model(xc, t, context=None, y=s, mask=mask)
|
||||
elif self.conditioning_key == 'hybrid-adm-mask':
|
||||
cc = torch.cat(c_crossattn, 1)
|
||||
if c_concat is not None:
|
||||
xc = torch.cat([x] + c_concat, dim=1)
|
||||
else:
|
||||
xc = x
|
||||
out = self.diffusion_model(xc, t, context=cc, y=s, mask=mask)
|
||||
elif self.conditioning_key == 'hybrid-time-adm': # adm means y, e.g., class index
|
||||
# assert s is not None
|
||||
assert c_adm is not None
|
||||
xc = torch.cat([x] + c_concat, dim=1)
|
||||
cc = torch.cat(c_crossattn, 1)
|
||||
out = self.diffusion_model(xc, t, context=cc, s=s, y=c_adm)
|
||||
elif self.conditioning_key == 'crossattn-adm':
|
||||
assert c_adm is not None
|
||||
cc = torch.cat(c_crossattn, 1)
|
||||
out = self.diffusion_model(x, t, context=cc, y=c_adm)
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,317 @@
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
import torch
|
||||
from ..models.utils_diffusion import make_ddim_sampling_parameters, make_ddim_timesteps, rescale_noise_cfg
|
||||
from ..common import noise_like
|
||||
from ..common import extract_into_tensor
|
||||
import copy
|
||||
|
||||
|
||||
class DDIMSampler(object):
|
||||
def __init__(self, model, schedule="linear", **kwargs):
|
||||
super().__init__()
|
||||
self.model = model
|
||||
self.ddpm_num_timesteps = model.num_timesteps
|
||||
self.schedule = schedule
|
||||
self.counter = 0
|
||||
|
||||
def register_buffer(self, name, attr):
|
||||
if type(attr) == torch.Tensor:
|
||||
if attr.device != torch.device("cuda"):
|
||||
attr = attr.to(torch.device("cuda"))
|
||||
setattr(self, name, attr)
|
||||
|
||||
def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., verbose=True):
|
||||
self.ddim_timesteps = make_ddim_timesteps(ddim_discr_method=ddim_discretize, num_ddim_timesteps=ddim_num_steps,
|
||||
num_ddpm_timesteps=self.ddpm_num_timesteps,verbose=verbose)
|
||||
alphas_cumprod = self.model.alphas_cumprod
|
||||
assert alphas_cumprod.shape[0] == self.ddpm_num_timesteps, 'alphas have to be defined for each timestep'
|
||||
to_torch = lambda x: x.clone().detach().to(torch.float32).to(self.model.device)
|
||||
|
||||
if self.model.use_dynamic_rescale:
|
||||
self.ddim_scale_arr = self.model.scale_arr[self.ddim_timesteps]
|
||||
self.ddim_scale_arr_prev = torch.cat([self.ddim_scale_arr[0:1], self.ddim_scale_arr[:-1]])
|
||||
|
||||
self.register_buffer('betas', to_torch(self.model.betas))
|
||||
self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod))
|
||||
self.register_buffer('alphas_cumprod_prev', to_torch(self.model.alphas_cumprod_prev))
|
||||
|
||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
||||
self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod.cpu())))
|
||||
self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod.cpu())))
|
||||
self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod.cpu())))
|
||||
self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu())))
|
||||
self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu() - 1)))
|
||||
|
||||
# ddim sampling parameters
|
||||
ddim_sigmas, ddim_alphas, ddim_alphas_prev = make_ddim_sampling_parameters(alphacums=alphas_cumprod.cpu(),
|
||||
ddim_timesteps=self.ddim_timesteps,
|
||||
eta=ddim_eta,verbose=verbose)
|
||||
self.register_buffer('ddim_sigmas', ddim_sigmas)
|
||||
self.register_buffer('ddim_alphas', ddim_alphas)
|
||||
self.register_buffer('ddim_alphas_prev', ddim_alphas_prev)
|
||||
self.register_buffer('ddim_sqrt_one_minus_alphas', np.sqrt(1. - ddim_alphas))
|
||||
sigmas_for_original_sampling_steps = ddim_eta * torch.sqrt(
|
||||
(1 - self.alphas_cumprod_prev) / (1 - self.alphas_cumprod) * (
|
||||
1 - self.alphas_cumprod / self.alphas_cumprod_prev))
|
||||
self.register_buffer('ddim_sigmas_for_original_num_steps', sigmas_for_original_sampling_steps)
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self,
|
||||
S,
|
||||
batch_size,
|
||||
shape,
|
||||
conditioning=None,
|
||||
callback=None,
|
||||
normals_sequence=None,
|
||||
img_callback=None,
|
||||
quantize_x0=False,
|
||||
eta=0.,
|
||||
mask=None,
|
||||
x0=None,
|
||||
temperature=1.,
|
||||
noise_dropout=0.,
|
||||
score_corrector=None,
|
||||
corrector_kwargs=None,
|
||||
verbose=True,
|
||||
schedule_verbose=False,
|
||||
x_T=None,
|
||||
log_every_t=100,
|
||||
unconditional_guidance_scale=1.,
|
||||
unconditional_conditioning=None,
|
||||
precision=None,
|
||||
fs=None,
|
||||
timestep_spacing='uniform', #uniform_trailing for starting from last timestep
|
||||
guidance_rescale=0.0,
|
||||
**kwargs
|
||||
):
|
||||
|
||||
# check condition bs
|
||||
if conditioning is not None:
|
||||
if isinstance(conditioning, dict):
|
||||
try:
|
||||
cbs = conditioning[list(conditioning.keys())[0]].shape[0]
|
||||
except:
|
||||
cbs = conditioning[list(conditioning.keys())[0]][0].shape[0]
|
||||
|
||||
if cbs != batch_size:
|
||||
print(f"Warning: Got {cbs} conditionings but batch-size is {batch_size}")
|
||||
else:
|
||||
if conditioning.shape[0] != batch_size:
|
||||
print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}")
|
||||
|
||||
self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, verbose=schedule_verbose)
|
||||
|
||||
# make shape
|
||||
if len(shape) == 3:
|
||||
C, H, W = shape
|
||||
size = (batch_size, C, H, W)
|
||||
elif len(shape) == 4:
|
||||
C, T, H, W = shape
|
||||
size = (batch_size, C, T, H, W)
|
||||
|
||||
samples, intermediates = self.ddim_sampling(conditioning, size,
|
||||
callback=callback,
|
||||
img_callback=img_callback,
|
||||
quantize_denoised=quantize_x0,
|
||||
mask=mask, x0=x0,
|
||||
ddim_use_original_steps=False,
|
||||
noise_dropout=noise_dropout,
|
||||
temperature=temperature,
|
||||
score_corrector=score_corrector,
|
||||
corrector_kwargs=corrector_kwargs,
|
||||
x_T=x_T,
|
||||
log_every_t=log_every_t,
|
||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
||||
unconditional_conditioning=unconditional_conditioning,
|
||||
verbose=verbose,
|
||||
precision=precision,
|
||||
fs=fs,
|
||||
guidance_rescale=guidance_rescale,
|
||||
**kwargs)
|
||||
return samples, intermediates
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_sampling(self, cond, shape,
|
||||
x_T=None, ddim_use_original_steps=False,
|
||||
callback=None, timesteps=None, quantize_denoised=False,
|
||||
mask=None, x0=None, img_callback=None, log_every_t=100,
|
||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
||||
unconditional_guidance_scale=1., unconditional_conditioning=None, verbose=True,precision=None,fs=None,guidance_rescale=0.0,
|
||||
**kwargs):
|
||||
device = self.model.betas.device
|
||||
b = shape[0]
|
||||
if x_T is None:
|
||||
img = torch.randn(shape, device=device)
|
||||
else:
|
||||
img = x_T
|
||||
if precision is not None:
|
||||
if precision == 16:
|
||||
img = img.to(dtype=torch.float16)
|
||||
|
||||
if timesteps is None:
|
||||
timesteps = self.ddpm_num_timesteps if ddim_use_original_steps else self.ddim_timesteps
|
||||
elif timesteps is not None and not ddim_use_original_steps:
|
||||
subset_end = int(min(timesteps / self.ddim_timesteps.shape[0], 1) * self.ddim_timesteps.shape[0]) - 1
|
||||
timesteps = self.ddim_timesteps[:subset_end]
|
||||
|
||||
intermediates = {'x_inter': [img], 'pred_x0': [img]}
|
||||
time_range = reversed(range(0,timesteps)) if ddim_use_original_steps else np.flip(timesteps)
|
||||
total_steps = timesteps if ddim_use_original_steps else timesteps.shape[0]
|
||||
if verbose:
|
||||
iterator = tqdm(time_range, desc='DDIM Sampler', total=total_steps)
|
||||
else:
|
||||
iterator = time_range
|
||||
|
||||
clean_cond = kwargs.pop("clean_cond", False)
|
||||
|
||||
# cond_copy, unconditional_conditioning_copy = copy.deepcopy(cond), copy.deepcopy(unconditional_conditioning)
|
||||
for i, step in enumerate(iterator):
|
||||
index = total_steps - i - 1
|
||||
ts = torch.full((b,), step, device=device, dtype=torch.long)
|
||||
|
||||
## use mask to blend noised original latent (img_orig) & new sampled latent (img)
|
||||
if mask is not None:
|
||||
assert x0 is not None
|
||||
if clean_cond:
|
||||
img_orig = x0
|
||||
else:
|
||||
img_orig = self.model.q_sample(x0, ts) # TODO: deterministic forward pass? <ddim inversion>
|
||||
img = img_orig * mask + (1. - mask) * img # keep original & modify use img
|
||||
|
||||
|
||||
|
||||
|
||||
outs = self.p_sample_ddim(img, cond, ts, index=index, use_original_steps=ddim_use_original_steps,
|
||||
quantize_denoised=quantize_denoised, temperature=temperature,
|
||||
noise_dropout=noise_dropout, score_corrector=score_corrector,
|
||||
corrector_kwargs=corrector_kwargs,
|
||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
||||
unconditional_conditioning=unconditional_conditioning,
|
||||
mask=mask,x0=x0,fs=fs,guidance_rescale=guidance_rescale,
|
||||
**kwargs)
|
||||
|
||||
|
||||
img, pred_x0 = outs
|
||||
if callback: callback(i)
|
||||
if img_callback: img_callback(pred_x0, i)
|
||||
|
||||
if index % log_every_t == 0 or index == total_steps - 1:
|
||||
intermediates['x_inter'].append(img)
|
||||
intermediates['pred_x0'].append(pred_x0)
|
||||
|
||||
return img, intermediates
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample_ddim(self, x, c, t, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False,
|
||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
||||
unconditional_guidance_scale=1., unconditional_conditioning=None,
|
||||
uc_type=None, conditional_guidance_scale_temporal=None,mask=None,x0=None,guidance_rescale=0.0,**kwargs):
|
||||
b, *_, device = *x.shape, x.device
|
||||
if x.dim() == 5:
|
||||
is_video = True
|
||||
else:
|
||||
is_video = False
|
||||
|
||||
if unconditional_conditioning is None or unconditional_guidance_scale == 1.:
|
||||
model_output = self.model.apply_model(x, t, c, **kwargs) # unet denoiser
|
||||
else:
|
||||
### do_classifier_free_guidance
|
||||
if isinstance(c, torch.Tensor) or isinstance(c, dict):
|
||||
e_t_cond = self.model.apply_model(x, t, c, **kwargs)
|
||||
e_t_uncond = self.model.apply_model(x, t, unconditional_conditioning, **kwargs)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
model_output = e_t_uncond + unconditional_guidance_scale * (e_t_cond - e_t_uncond)
|
||||
|
||||
if guidance_rescale > 0.0:
|
||||
model_output = rescale_noise_cfg(model_output, e_t_cond, guidance_rescale=guidance_rescale)
|
||||
|
||||
if self.model.parameterization == "v":
|
||||
e_t = self.model.predict_eps_from_z_and_v(x, t, model_output)
|
||||
else:
|
||||
e_t = model_output
|
||||
|
||||
if score_corrector is not None:
|
||||
assert self.model.parameterization == "eps", 'not implemented'
|
||||
e_t = score_corrector.modify_score(self.model, e_t, x, t, c, **corrector_kwargs)
|
||||
|
||||
alphas = self.model.alphas_cumprod if use_original_steps else self.ddim_alphas
|
||||
alphas_prev = self.model.alphas_cumprod_prev if use_original_steps else self.ddim_alphas_prev
|
||||
sqrt_one_minus_alphas = self.model.sqrt_one_minus_alphas_cumprod if use_original_steps else self.ddim_sqrt_one_minus_alphas
|
||||
# sigmas = self.model.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
|
||||
sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
|
||||
# select parameters corresponding to the currently considered timestep
|
||||
|
||||
if is_video:
|
||||
size = (b, 1, 1, 1, 1)
|
||||
else:
|
||||
size = (b, 1, 1, 1)
|
||||
a_t = torch.full(size, alphas[index], device=device)
|
||||
a_prev = torch.full(size, alphas_prev[index], device=device)
|
||||
sigma_t = torch.full(size, sigmas[index], device=device)
|
||||
sqrt_one_minus_at = torch.full(size, sqrt_one_minus_alphas[index],device=device)
|
||||
|
||||
# current prediction for x_0
|
||||
if self.model.parameterization != "v":
|
||||
pred_x0 = (x - sqrt_one_minus_at * e_t) / a_t.sqrt()
|
||||
else:
|
||||
pred_x0 = self.model.predict_start_from_z_and_v(x, t, model_output)
|
||||
|
||||
if self.model.use_dynamic_rescale:
|
||||
scale_t = torch.full(size, self.ddim_scale_arr[index], device=device)
|
||||
prev_scale_t = torch.full(size, self.ddim_scale_arr_prev[index], device=device)
|
||||
rescale = (prev_scale_t / scale_t)
|
||||
pred_x0 *= rescale
|
||||
|
||||
if quantize_denoised:
|
||||
pred_x0, _, *_ = self.model.first_stage_model.quantize(pred_x0)
|
||||
# direction pointing to x_t
|
||||
dir_xt = (1. - a_prev - sigma_t**2).sqrt() * e_t
|
||||
|
||||
noise = sigma_t * noise_like(x.shape, device, repeat_noise) * temperature
|
||||
if noise_dropout > 0.:
|
||||
noise = torch.nn.functional.dropout(noise, p=noise_dropout)
|
||||
|
||||
x_prev = a_prev.sqrt() * pred_x0 + dir_xt + noise
|
||||
|
||||
return x_prev, pred_x0
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, x_latent, cond, t_start, unconditional_guidance_scale=1.0, unconditional_conditioning=None,
|
||||
use_original_steps=False, callback=None):
|
||||
|
||||
timesteps = np.arange(self.ddpm_num_timesteps) if use_original_steps else self.ddim_timesteps
|
||||
timesteps = timesteps[:t_start]
|
||||
|
||||
time_range = np.flip(timesteps)
|
||||
total_steps = timesteps.shape[0]
|
||||
print(f"Running DDIM Sampling with {total_steps} timesteps")
|
||||
|
||||
iterator = tqdm(time_range, desc='Decoding image', total=total_steps)
|
||||
x_dec = x_latent
|
||||
for i, step in enumerate(iterator):
|
||||
index = total_steps - i - 1
|
||||
ts = torch.full((x_latent.shape[0],), step, device=x_latent.device, dtype=torch.long)
|
||||
x_dec, _ = self.p_sample_ddim(x_dec, cond, ts, index=index, use_original_steps=use_original_steps,
|
||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
||||
unconditional_conditioning=unconditional_conditioning)
|
||||
if callback: callback(i)
|
||||
return x_dec
|
||||
|
||||
@torch.no_grad()
|
||||
def stochastic_encode(self, x0, t, use_original_steps=False, noise=None):
|
||||
# fast, but does not allow for exact reconstruction
|
||||
# t serves as an index to gather the correct alphas
|
||||
if use_original_steps:
|
||||
sqrt_alphas_cumprod = self.sqrt_alphas_cumprod
|
||||
sqrt_one_minus_alphas_cumprod = self.sqrt_one_minus_alphas_cumprod
|
||||
else:
|
||||
sqrt_alphas_cumprod = torch.sqrt(self.ddim_alphas)
|
||||
sqrt_one_minus_alphas_cumprod = self.ddim_sqrt_one_minus_alphas
|
||||
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x0)
|
||||
return (extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 +
|
||||
extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) * noise)
|
||||
@@ -0,0 +1,323 @@
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
import torch
|
||||
from ...models.utils_diffusion import make_ddim_sampling_parameters, make_ddim_timesteps, rescale_noise_cfg
|
||||
from ..common import noise_like
|
||||
from ..common import extract_into_tensor
|
||||
import copy
|
||||
|
||||
|
||||
class DDIMSampler(object):
|
||||
def __init__(self, model, schedule="linear", **kwargs):
|
||||
super().__init__()
|
||||
self.model = model
|
||||
self.ddpm_num_timesteps = model.num_timesteps
|
||||
self.schedule = schedule
|
||||
self.counter = 0
|
||||
|
||||
def register_buffer(self, name, attr):
|
||||
if type(attr) == torch.Tensor:
|
||||
if attr.device != torch.device("cuda"):
|
||||
attr = attr.to(torch.device("cuda"))
|
||||
setattr(self, name, attr)
|
||||
|
||||
def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., verbose=True):
|
||||
self.ddim_timesteps = make_ddim_timesteps(ddim_discr_method=ddim_discretize, num_ddim_timesteps=ddim_num_steps,
|
||||
num_ddpm_timesteps=self.ddpm_num_timesteps,verbose=verbose)
|
||||
alphas_cumprod = self.model.alphas_cumprod
|
||||
assert alphas_cumprod.shape[0] == self.ddpm_num_timesteps, 'alphas have to be defined for each timestep'
|
||||
to_torch = lambda x: x.clone().detach().to(torch.float32).to(self.model.device)
|
||||
|
||||
if self.model.use_dynamic_rescale:
|
||||
self.ddim_scale_arr = self.model.scale_arr[self.ddim_timesteps]
|
||||
self.ddim_scale_arr_prev = torch.cat([self.ddim_scale_arr[0:1], self.ddim_scale_arr[:-1]])
|
||||
|
||||
self.register_buffer('betas', to_torch(self.model.betas))
|
||||
self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod))
|
||||
self.register_buffer('alphas_cumprod_prev', to_torch(self.model.alphas_cumprod_prev))
|
||||
|
||||
# calculations for diffusion q(x_t | x_{t-1}) and others
|
||||
self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod.cpu())))
|
||||
self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod.cpu())))
|
||||
self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod.cpu())))
|
||||
self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu())))
|
||||
self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu() - 1)))
|
||||
|
||||
# ddim sampling parameters
|
||||
ddim_sigmas, ddim_alphas, ddim_alphas_prev = make_ddim_sampling_parameters(alphacums=alphas_cumprod.cpu(),
|
||||
ddim_timesteps=self.ddim_timesteps,
|
||||
eta=ddim_eta,verbose=verbose)
|
||||
self.register_buffer('ddim_sigmas', ddim_sigmas)
|
||||
self.register_buffer('ddim_alphas', ddim_alphas)
|
||||
self.register_buffer('ddim_alphas_prev', ddim_alphas_prev)
|
||||
self.register_buffer('ddim_sqrt_one_minus_alphas', np.sqrt(1. - ddim_alphas))
|
||||
sigmas_for_original_sampling_steps = ddim_eta * torch.sqrt(
|
||||
(1 - self.alphas_cumprod_prev) / (1 - self.alphas_cumprod) * (
|
||||
1 - self.alphas_cumprod / self.alphas_cumprod_prev))
|
||||
self.register_buffer('ddim_sigmas_for_original_num_steps', sigmas_for_original_sampling_steps)
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self,
|
||||
S,
|
||||
batch_size,
|
||||
shape,
|
||||
conditioning=None,
|
||||
callback=None,
|
||||
normals_sequence=None,
|
||||
img_callback=None,
|
||||
quantize_x0=False,
|
||||
eta=0.,
|
||||
mask=None,
|
||||
x0=None,
|
||||
temperature=1.,
|
||||
noise_dropout=0.,
|
||||
score_corrector=None,
|
||||
corrector_kwargs=None,
|
||||
verbose=True,
|
||||
schedule_verbose=False,
|
||||
x_T=None,
|
||||
log_every_t=100,
|
||||
unconditional_guidance_scale=1.,
|
||||
unconditional_conditioning=None,
|
||||
precision=None,
|
||||
fs=None,
|
||||
timestep_spacing='uniform', #uniform_trailing for starting from last timestep
|
||||
guidance_rescale=0.0,
|
||||
# this has to come in the same format as the conditioning, # e.g. as encoded tokens, ...
|
||||
**kwargs
|
||||
):
|
||||
|
||||
# check condition bs
|
||||
if conditioning is not None:
|
||||
if isinstance(conditioning, dict):
|
||||
try:
|
||||
cbs = conditioning[list(conditioning.keys())[0]].shape[0]
|
||||
except:
|
||||
cbs = conditioning[list(conditioning.keys())[0]][0].shape[0]
|
||||
|
||||
if cbs != batch_size:
|
||||
print(f"Warning: Got {cbs} conditionings but batch-size is {batch_size}")
|
||||
else:
|
||||
if conditioning.shape[0] != batch_size:
|
||||
print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}")
|
||||
|
||||
# print('==> timestep_spacing: ', timestep_spacing, guidance_rescale)
|
||||
self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, verbose=schedule_verbose)
|
||||
|
||||
# make shape
|
||||
if len(shape) == 3:
|
||||
C, H, W = shape
|
||||
size = (batch_size, C, H, W)
|
||||
elif len(shape) == 4:
|
||||
C, T, H, W = shape
|
||||
size = (batch_size, C, T, H, W)
|
||||
# print(f'Data shape for DDIM sampling is {size}, eta {eta}')
|
||||
|
||||
samples, intermediates = self.ddim_sampling(conditioning, size,
|
||||
callback=callback,
|
||||
img_callback=img_callback,
|
||||
quantize_denoised=quantize_x0,
|
||||
mask=mask, x0=x0,
|
||||
ddim_use_original_steps=False,
|
||||
noise_dropout=noise_dropout,
|
||||
temperature=temperature,
|
||||
score_corrector=score_corrector,
|
||||
corrector_kwargs=corrector_kwargs,
|
||||
x_T=x_T,
|
||||
log_every_t=log_every_t,
|
||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
||||
unconditional_conditioning=unconditional_conditioning,
|
||||
verbose=verbose,
|
||||
precision=precision,
|
||||
fs=fs,
|
||||
guidance_rescale=guidance_rescale,
|
||||
**kwargs)
|
||||
return samples, intermediates
|
||||
|
||||
@torch.no_grad()
|
||||
def ddim_sampling(self, cond, shape,
|
||||
x_T=None, ddim_use_original_steps=False,
|
||||
callback=None, timesteps=None, quantize_denoised=False,
|
||||
mask=None, x0=None, img_callback=None, log_every_t=100,
|
||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
||||
unconditional_guidance_scale=1., unconditional_conditioning=None, verbose=True,precision=None,fs=None,guidance_rescale=0.0,
|
||||
**kwargs):
|
||||
device = self.model.betas.device
|
||||
b = shape[0]
|
||||
if x_T is None:
|
||||
img = torch.randn(shape, device=device)
|
||||
else:
|
||||
img = x_T
|
||||
if precision is not None:
|
||||
if precision == 16:
|
||||
img = img.to(dtype=torch.float16)
|
||||
|
||||
|
||||
if timesteps is None:
|
||||
timesteps = self.ddpm_num_timesteps if ddim_use_original_steps else self.ddim_timesteps
|
||||
elif timesteps is not None and not ddim_use_original_steps:
|
||||
subset_end = int(min(timesteps / self.ddim_timesteps.shape[0], 1) * self.ddim_timesteps.shape[0]) - 1
|
||||
timesteps = self.ddim_timesteps[:subset_end]
|
||||
|
||||
intermediates = {'x_inter': [img], 'pred_x0': [img]}
|
||||
time_range = reversed(range(0,timesteps)) if ddim_use_original_steps else np.flip(timesteps)
|
||||
total_steps = timesteps if ddim_use_original_steps else timesteps.shape[0]
|
||||
if verbose:
|
||||
iterator = tqdm(time_range, desc='DDIM Sampler', total=total_steps)
|
||||
else:
|
||||
iterator = time_range
|
||||
|
||||
clean_cond = kwargs.pop("clean_cond", False)
|
||||
|
||||
# cond_copy, unconditional_conditioning_copy = copy.deepcopy(cond), copy.deepcopy(unconditional_conditioning)
|
||||
for i, step in enumerate(iterator):
|
||||
index = total_steps - i - 1
|
||||
ts = torch.full((b,), step, device=device, dtype=torch.long)
|
||||
|
||||
## use mask to blend noised original latent (img_orig) & new sampled latent (img)
|
||||
if mask is not None:
|
||||
assert x0 is not None
|
||||
if clean_cond:
|
||||
img_orig = x0
|
||||
else:
|
||||
img_orig = self.model.q_sample(x0, ts) # TODO: deterministic forward pass? <ddim inversion>
|
||||
img = img_orig * mask + (1. - mask) * img # keep original & modify use img
|
||||
|
||||
|
||||
|
||||
|
||||
outs = self.p_sample_ddim(img, cond, ts, index=index, use_original_steps=ddim_use_original_steps,
|
||||
quantize_denoised=quantize_denoised, temperature=temperature,
|
||||
noise_dropout=noise_dropout, score_corrector=score_corrector,
|
||||
corrector_kwargs=corrector_kwargs,
|
||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
||||
unconditional_conditioning=unconditional_conditioning,
|
||||
mask=mask,x0=x0,fs=fs,guidance_rescale=guidance_rescale,
|
||||
**kwargs)
|
||||
|
||||
|
||||
|
||||
img, pred_x0 = outs
|
||||
if callback: callback(i)
|
||||
if img_callback: img_callback(pred_x0, i)
|
||||
|
||||
if index % log_every_t == 0 or index == total_steps - 1:
|
||||
intermediates['x_inter'].append(img)
|
||||
intermediates['pred_x0'].append(pred_x0)
|
||||
|
||||
return img, intermediates
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample_ddim(self, x, c, t, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False,
|
||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
||||
unconditional_guidance_scale=1., unconditional_conditioning=None,
|
||||
uc_type=None, cfg_img=None,mask=None,x0=None,guidance_rescale=0.0, **kwargs):
|
||||
b, *_, device = *x.shape, x.device
|
||||
if x.dim() == 5:
|
||||
is_video = True
|
||||
else:
|
||||
is_video = False
|
||||
if cfg_img is None:
|
||||
cfg_img = unconditional_guidance_scale
|
||||
|
||||
unconditional_conditioning_img_nonetext = kwargs['unconditional_conditioning_img_nonetext']
|
||||
|
||||
|
||||
if unconditional_conditioning is None or unconditional_guidance_scale == 1.:
|
||||
model_output = self.model.apply_model(x, t, c, **kwargs) # unet denoiser
|
||||
else:
|
||||
### with unconditional condition
|
||||
e_t_cond = self.model.apply_model(x, t, c, **kwargs)
|
||||
e_t_uncond = self.model.apply_model(x, t, unconditional_conditioning, **kwargs)
|
||||
e_t_uncond_img = self.model.apply_model(x, t, unconditional_conditioning_img_nonetext, **kwargs)
|
||||
# text cfg
|
||||
model_output = e_t_uncond + cfg_img * (e_t_uncond_img - e_t_uncond) + unconditional_guidance_scale * (e_t_cond - e_t_uncond_img)
|
||||
if guidance_rescale > 0.0:
|
||||
model_output = rescale_noise_cfg(model_output, e_t_cond, guidance_rescale=guidance_rescale)
|
||||
|
||||
if self.model.parameterization == "v":
|
||||
e_t = self.model.predict_eps_from_z_and_v(x, t, model_output)
|
||||
else:
|
||||
e_t = model_output
|
||||
|
||||
if score_corrector is not None:
|
||||
assert self.model.parameterization == "eps", 'not implemented'
|
||||
e_t = score_corrector.modify_score(self.model, e_t, x, t, c, **corrector_kwargs)
|
||||
|
||||
alphas = self.model.alphas_cumprod if use_original_steps else self.ddim_alphas
|
||||
alphas_prev = self.model.alphas_cumprod_prev if use_original_steps else self.ddim_alphas_prev
|
||||
sqrt_one_minus_alphas = self.model.sqrt_one_minus_alphas_cumprod if use_original_steps else self.ddim_sqrt_one_minus_alphas
|
||||
sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
|
||||
# select parameters corresponding to the currently considered timestep
|
||||
|
||||
if is_video:
|
||||
size = (b, 1, 1, 1, 1)
|
||||
else:
|
||||
size = (b, 1, 1, 1)
|
||||
a_t = torch.full(size, alphas[index], device=device)
|
||||
a_prev = torch.full(size, alphas_prev[index], device=device)
|
||||
sigma_t = torch.full(size, sigmas[index], device=device)
|
||||
sqrt_one_minus_at = torch.full(size, sqrt_one_minus_alphas[index],device=device)
|
||||
|
||||
# current prediction for x_0
|
||||
if self.model.parameterization != "v":
|
||||
pred_x0 = (x - sqrt_one_minus_at * e_t) / a_t.sqrt()
|
||||
else:
|
||||
pred_x0 = self.model.predict_start_from_z_and_v(x, t, model_output)
|
||||
|
||||
if self.model.use_dynamic_rescale:
|
||||
scale_t = torch.full(size, self.ddim_scale_arr[index], device=device)
|
||||
prev_scale_t = torch.full(size, self.ddim_scale_arr_prev[index], device=device)
|
||||
rescale = (prev_scale_t / scale_t)
|
||||
pred_x0 *= rescale
|
||||
|
||||
if quantize_denoised:
|
||||
pred_x0, _, *_ = self.model.first_stage_model.quantize(pred_x0)
|
||||
# direction pointing to x_t
|
||||
dir_xt = (1. - a_prev - sigma_t**2).sqrt() * e_t
|
||||
|
||||
noise = sigma_t * noise_like(x.shape, device, repeat_noise) * temperature
|
||||
if noise_dropout > 0.:
|
||||
noise = torch.nn.functional.dropout(noise, p=noise_dropout)
|
||||
|
||||
x_prev = a_prev.sqrt() * pred_x0 + dir_xt + noise
|
||||
|
||||
return x_prev, pred_x0
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, x_latent, cond, t_start, unconditional_guidance_scale=1.0, unconditional_conditioning=None,
|
||||
use_original_steps=False, callback=None):
|
||||
|
||||
timesteps = np.arange(self.ddpm_num_timesteps) if use_original_steps else self.ddim_timesteps
|
||||
timesteps = timesteps[:t_start]
|
||||
|
||||
time_range = np.flip(timesteps)
|
||||
total_steps = timesteps.shape[0]
|
||||
print(f"Running DDIM Sampling with {total_steps} timesteps")
|
||||
|
||||
iterator = tqdm(time_range, desc='Decoding image', total=total_steps)
|
||||
x_dec = x_latent
|
||||
for i, step in enumerate(iterator):
|
||||
index = total_steps - i - 1
|
||||
ts = torch.full((x_latent.shape[0],), step, device=x_latent.device, dtype=torch.long)
|
||||
x_dec, _ = self.p_sample_ddim(x_dec, cond, ts, index=index, use_original_steps=use_original_steps,
|
||||
unconditional_guidance_scale=unconditional_guidance_scale,
|
||||
unconditional_conditioning=unconditional_conditioning)
|
||||
if callback: callback(i)
|
||||
return x_dec
|
||||
|
||||
@torch.no_grad()
|
||||
def stochastic_encode(self, x0, t, use_original_steps=False, noise=None):
|
||||
# fast, but does not allow for exact reconstruction
|
||||
# t serves as an index to gather the correct alphas
|
||||
if use_original_steps:
|
||||
sqrt_alphas_cumprod = self.sqrt_alphas_cumprod
|
||||
sqrt_one_minus_alphas_cumprod = self.sqrt_one_minus_alphas_cumprod
|
||||
else:
|
||||
sqrt_alphas_cumprod = torch.sqrt(self.ddim_alphas)
|
||||
sqrt_one_minus_alphas_cumprod = self.ddim_sqrt_one_minus_alphas
|
||||
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x0)
|
||||
return (extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 +
|
||||
extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) * noise)
|
||||
@@ -0,0 +1 @@
|
||||
from .sampler import UniPCSampler
|
||||
@@ -0,0 +1,79 @@
|
||||
"""SAMPLING ONLY."""
|
||||
|
||||
import torch
|
||||
|
||||
from .uni_pc import NoiseScheduleVP, model_wrapper, UniPC
|
||||
|
||||
class UniPCSampler(object):
|
||||
def __init__(self, model, **kwargs):
|
||||
super().__init__()
|
||||
self.model = model
|
||||
to_torch = lambda x: x.clone().detach().to(torch.float32).to(model.device)
|
||||
self.register_buffer('alphas_cumprod', to_torch(model.alphas_cumprod))
|
||||
|
||||
def register_buffer(self, name, attr):
|
||||
if type(attr) == torch.Tensor:
|
||||
if attr.device != torch.device("cuda"):
|
||||
attr = attr.to(torch.device("cuda"))
|
||||
setattr(self, name, attr)
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self,
|
||||
S,
|
||||
batch_size,
|
||||
shape,
|
||||
conditioning=None,
|
||||
callback=None,
|
||||
normals_sequence=None,
|
||||
img_callback=None,
|
||||
quantize_x0=False,
|
||||
eta=0.,
|
||||
mask=None,
|
||||
x0=None,
|
||||
temperature=1.,
|
||||
noise_dropout=0.,
|
||||
score_corrector=None,
|
||||
corrector_kwargs=None,
|
||||
verbose=True,
|
||||
x_T=None,
|
||||
log_every_t=100,
|
||||
unconditional_guidance_scale=1.,
|
||||
unconditional_conditioning=None,
|
||||
# this has to come in the same format as the conditioning, # e.g. as encoded tokens, ...
|
||||
**kwargs
|
||||
):
|
||||
if conditioning is not None:
|
||||
if isinstance(conditioning, dict):
|
||||
cbs = conditioning[list(conditioning.keys())[0]].shape[0]
|
||||
if cbs != batch_size:
|
||||
print(f"Warning: Got {cbs} conditionings but batch-size is {batch_size}")
|
||||
else:
|
||||
if conditioning.shape[0] != batch_size:
|
||||
print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}")
|
||||
|
||||
# sampling
|
||||
C, F, H, W = shape
|
||||
size = (batch_size, C, H, W)
|
||||
|
||||
device = self.model.betas.device
|
||||
if x_T is None:
|
||||
img = torch.randn(size, device=device)
|
||||
else:
|
||||
img = x_T
|
||||
|
||||
ns = NoiseScheduleVP('discrete', alphas_cumprod=self.alphas_cumprod)
|
||||
|
||||
model_fn = model_wrapper(
|
||||
lambda x, t, c: self.model.apply_model(x, t, c),
|
||||
ns,
|
||||
model_type="noise",
|
||||
guidance_type="classifier-free",
|
||||
condition=conditioning,
|
||||
unconditional_condition=unconditional_conditioning,
|
||||
guidance_scale=unconditional_guidance_scale,
|
||||
)
|
||||
|
||||
uni_pc = UniPC(model_fn, ns, predict_x0=True, thresholding=False)
|
||||
x = uni_pc.sample(img, steps=S, skip_type="time_uniform", method="multistep", order=3, lower_order_final=True)
|
||||
|
||||
return x.to(device), None
|
||||
@@ -0,0 +1,808 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import math
|
||||
|
||||
|
||||
class NoiseScheduleVP:
|
||||
def __init__(
|
||||
self,
|
||||
schedule='discrete',
|
||||
betas=None,
|
||||
alphas_cumprod=None,
|
||||
continuous_beta_0=0.1,
|
||||
continuous_beta_1=20.,
|
||||
):
|
||||
"""Create a wrapper class for the forward SDE (VP type).
|
||||
|
||||
***
|
||||
Update: We support discrete-time diffusion models by implementing a picewise linear interpolation for log_alpha_t.
|
||||
We recommend to use schedule='discrete' for the discrete-time diffusion models, especially for high-resolution images.
|
||||
***
|
||||
|
||||
The forward SDE ensures that the condition distribution q_{t|0}(x_t | x_0) = N ( alpha_t * x_0, sigma_t^2 * I ).
|
||||
We further define lambda_t = log(alpha_t) - log(sigma_t), which is the half-logSNR (described in the DPM-Solver paper).
|
||||
Therefore, we implement the functions for computing alpha_t, sigma_t and lambda_t. For t in [0, T], we have:
|
||||
|
||||
log_alpha_t = self.marginal_log_mean_coeff(t)
|
||||
sigma_t = self.marginal_std(t)
|
||||
lambda_t = self.marginal_lambda(t)
|
||||
|
||||
Moreover, as lambda(t) is an invertible function, we also support its inverse function:
|
||||
|
||||
t = self.inverse_lambda(lambda_t)
|
||||
|
||||
===============================================================
|
||||
|
||||
We support both discrete-time DPMs (trained on n = 0, 1, ..., N-1) and continuous-time DPMs (trained on t in [t_0, T]).
|
||||
|
||||
1. For discrete-time DPMs:
|
||||
|
||||
For discrete-time DPMs trained on n = 0, 1, ..., N-1, we convert the discrete steps to continuous time steps by:
|
||||
t_i = (i + 1) / N
|
||||
e.g. for N = 1000, we have t_0 = 1e-3 and T = t_{N-1} = 1.
|
||||
We solve the corresponding diffusion ODE from time T = 1 to time t_0 = 1e-3.
|
||||
|
||||
Args:
|
||||
betas: A `torch.Tensor`. The beta array for the discrete-time DPM. (See the original DDPM paper for details)
|
||||
alphas_cumprod: A `torch.Tensor`. The cumprod alphas for the discrete-time DPM. (See the original DDPM paper for details)
|
||||
|
||||
Note that we always have alphas_cumprod = cumprod(betas). Therefore, we only need to set one of `betas` and `alphas_cumprod`.
|
||||
|
||||
**Important**: Please pay special attention for the args for `alphas_cumprod`:
|
||||
The `alphas_cumprod` is the \hat{alpha_n} arrays in the notations of DDPM. Specifically, DDPMs assume that
|
||||
q_{t_n | 0}(x_{t_n} | x_0) = N ( \sqrt{\hat{alpha_n}} * x_0, (1 - \hat{alpha_n}) * I ).
|
||||
Therefore, the notation \hat{alpha_n} is different from the notation alpha_t in DPM-Solver. In fact, we have
|
||||
alpha_{t_n} = \sqrt{\hat{alpha_n}},
|
||||
and
|
||||
log(alpha_{t_n}) = 0.5 * log(\hat{alpha_n}).
|
||||
|
||||
|
||||
2. For continuous-time DPMs:
|
||||
|
||||
We support two types of VPSDEs: linear (DDPM) and cosine (improved-DDPM). The hyperparameters for the noise
|
||||
schedule are the default settings in DDPM and improved-DDPM:
|
||||
|
||||
Args:
|
||||
beta_min: A `float` number. The smallest beta for the linear schedule.
|
||||
beta_max: A `float` number. The largest beta for the linear schedule.
|
||||
cosine_s: A `float` number. The hyperparameter in the cosine schedule.
|
||||
cosine_beta_max: A `float` number. The hyperparameter in the cosine schedule.
|
||||
T: A `float` number. The ending time of the forward process.
|
||||
|
||||
===============================================================
|
||||
|
||||
Args:
|
||||
schedule: A `str`. The noise schedule of the forward SDE. 'discrete' for discrete-time DPMs,
|
||||
'linear' or 'cosine' for continuous-time DPMs.
|
||||
Returns:
|
||||
A wrapper object of the forward SDE (VP type).
|
||||
|
||||
===============================================================
|
||||
|
||||
Example:
|
||||
|
||||
# For discrete-time DPMs, given betas (the beta array for n = 0, 1, ..., N - 1):
|
||||
>>> ns = NoiseScheduleVP('discrete', betas=betas)
|
||||
|
||||
# For discrete-time DPMs, given alphas_cumprod (the \hat{alpha_n} array for n = 0, 1, ..., N - 1):
|
||||
>>> ns = NoiseScheduleVP('discrete', alphas_cumprod=alphas_cumprod)
|
||||
|
||||
# For continuous-time DPMs (VPSDE), linear schedule:
|
||||
>>> ns = NoiseScheduleVP('linear', continuous_beta_0=0.1, continuous_beta_1=20.)
|
||||
|
||||
"""
|
||||
|
||||
if schedule not in ['discrete', 'linear', 'cosine']:
|
||||
raise ValueError("Unsupported noise schedule {}. The schedule needs to be 'discrete' or 'linear' or 'cosine'".format(schedule))
|
||||
|
||||
self.schedule = schedule
|
||||
if schedule == 'discrete':
|
||||
if betas is not None:
|
||||
log_alphas = 0.5 * torch.log(1 - betas).cumsum(dim=0)
|
||||
else:
|
||||
assert alphas_cumprod is not None
|
||||
log_alphas = 0.5 * torch.log(alphas_cumprod)
|
||||
self.total_N = len(log_alphas)
|
||||
self.T = 1.
|
||||
self.t_array = torch.linspace(0., 1., self.total_N + 1)[1:].reshape((1, -1))
|
||||
self.log_alpha_array = log_alphas.reshape((1, -1,))
|
||||
else:
|
||||
self.total_N = 1000
|
||||
self.beta_0 = continuous_beta_0
|
||||
self.beta_1 = continuous_beta_1
|
||||
self.cosine_s = 0.008
|
||||
self.cosine_beta_max = 999.
|
||||
self.cosine_t_max = math.atan(self.cosine_beta_max * (1. + self.cosine_s) / math.pi) * 2. * (1. + self.cosine_s) / math.pi - self.cosine_s
|
||||
self.cosine_log_alpha_0 = math.log(math.cos(self.cosine_s / (1. + self.cosine_s) * math.pi / 2.))
|
||||
self.schedule = schedule
|
||||
if schedule == 'cosine':
|
||||
# For the cosine schedule, T = 1 will have numerical issues. So we manually set the ending time T.
|
||||
# Note that T = 0.9946 may be not the optimal setting. However, we find it works well.
|
||||
self.T = 0.9946
|
||||
else:
|
||||
self.T = 1.
|
||||
|
||||
def marginal_log_mean_coeff(self, t):
|
||||
"""
|
||||
Compute log(alpha_t) of a given continuous-time label t in [0, T].
|
||||
"""
|
||||
if self.schedule == 'discrete':
|
||||
return interpolate_fn(t.reshape((-1, 1)), self.t_array.to(t.device), self.log_alpha_array.to(t.device)).reshape((-1))
|
||||
elif self.schedule == 'linear':
|
||||
return -0.25 * t ** 2 * (self.beta_1 - self.beta_0) - 0.5 * t * self.beta_0
|
||||
elif self.schedule == 'cosine':
|
||||
log_alpha_fn = lambda s: torch.log(torch.cos((s + self.cosine_s) / (1. + self.cosine_s) * math.pi / 2.))
|
||||
log_alpha_t = log_alpha_fn(t) - self.cosine_log_alpha_0
|
||||
return log_alpha_t
|
||||
|
||||
def marginal_alpha(self, t):
|
||||
"""
|
||||
Compute alpha_t of a given continuous-time label t in [0, T].
|
||||
"""
|
||||
return torch.exp(self.marginal_log_mean_coeff(t))
|
||||
|
||||
def marginal_std(self, t):
|
||||
"""
|
||||
Compute sigma_t of a given continuous-time label t in [0, T].
|
||||
"""
|
||||
return torch.sqrt(1. - torch.exp(2. * self.marginal_log_mean_coeff(t)))
|
||||
|
||||
def marginal_lambda(self, t):
|
||||
"""
|
||||
Compute lambda_t = log(alpha_t) - log(sigma_t) of a given continuous-time label t in [0, T].
|
||||
"""
|
||||
log_mean_coeff = self.marginal_log_mean_coeff(t)
|
||||
log_std = 0.5 * torch.log(1. - torch.exp(2. * log_mean_coeff))
|
||||
return log_mean_coeff - log_std
|
||||
|
||||
def inverse_lambda(self, lamb):
|
||||
"""
|
||||
Compute the continuous-time label t in [0, T] of a given half-logSNR lambda_t.
|
||||
"""
|
||||
if self.schedule == 'linear':
|
||||
tmp = 2. * (self.beta_1 - self.beta_0) * torch.logaddexp(-2. * lamb, torch.zeros((1,)).to(lamb))
|
||||
Delta = self.beta_0**2 + tmp
|
||||
return tmp / (torch.sqrt(Delta) + self.beta_0) / (self.beta_1 - self.beta_0)
|
||||
elif self.schedule == 'discrete':
|
||||
log_alpha = -0.5 * torch.logaddexp(torch.zeros((1,)).to(lamb.device), -2. * lamb)
|
||||
t = interpolate_fn(log_alpha.reshape((-1, 1)), torch.flip(self.log_alpha_array.to(lamb.device), [1]), torch.flip(self.t_array.to(lamb.device), [1]))
|
||||
return t.reshape((-1,))
|
||||
else:
|
||||
log_alpha = -0.5 * torch.logaddexp(-2. * lamb, torch.zeros((1,)).to(lamb))
|
||||
t_fn = lambda log_alpha_t: torch.arccos(torch.exp(log_alpha_t + self.cosine_log_alpha_0)) * 2. * (1. + self.cosine_s) / math.pi - self.cosine_s
|
||||
t = t_fn(log_alpha)
|
||||
return t
|
||||
|
||||
|
||||
def model_wrapper(
|
||||
model,
|
||||
noise_schedule,
|
||||
model_type="noise",
|
||||
model_kwargs={},
|
||||
guidance_type="uncond",
|
||||
condition=None,
|
||||
unconditional_condition=None,
|
||||
guidance_scale=1.,
|
||||
classifier_fn=None,
|
||||
classifier_kwargs={},
|
||||
):
|
||||
"""Create a wrapper function for the noise prediction model.
|
||||
|
||||
DPM-Solver needs to solve the continuous-time diffusion ODEs. For DPMs trained on discrete-time labels, we need to
|
||||
firstly wrap the model function to a noise prediction model that accepts the continuous time as the input.
|
||||
|
||||
We support four types of the diffusion model by setting `model_type`:
|
||||
|
||||
1. "noise": noise prediction model. (Trained by predicting noise).
|
||||
|
||||
2. "x_start": data prediction model. (Trained by predicting the data x_0 at time 0).
|
||||
|
||||
3. "v": velocity prediction model. (Trained by predicting the velocity).
|
||||
The "v" prediction is derivation detailed in Appendix D of [1], and is used in Imagen-Video [2].
|
||||
|
||||
[1] Salimans, Tim, and Jonathan Ho. "Progressive distillation for fast sampling of diffusion models."
|
||||
arXiv preprint arXiv:2202.00512 (2022).
|
||||
[2] Ho, Jonathan, et al. "Imagen Video: High Definition Video Generation with Diffusion Models."
|
||||
arXiv preprint arXiv:2210.02303 (2022).
|
||||
|
||||
4. "score": marginal score function. (Trained by denoising score matching).
|
||||
Note that the score function and the noise prediction model follows a simple relationship:
|
||||
```
|
||||
noise(x_t, t) = -sigma_t * score(x_t, t)
|
||||
```
|
||||
|
||||
We support three types of guided sampling by DPMs by setting `guidance_type`:
|
||||
1. "uncond": unconditional sampling by DPMs.
|
||||
The input `model` has the following format:
|
||||
``
|
||||
model(x, t_input, **model_kwargs) -> noise | x_start | v | score
|
||||
``
|
||||
|
||||
2. "classifier": classifier guidance sampling [3] by DPMs and another classifier.
|
||||
The input `model` has the following format:
|
||||
``
|
||||
model(x, t_input, **model_kwargs) -> noise | x_start | v | score
|
||||
``
|
||||
|
||||
The input `classifier_fn` has the following format:
|
||||
``
|
||||
classifier_fn(x, t_input, cond, **classifier_kwargs) -> logits(x, t_input, cond)
|
||||
``
|
||||
|
||||
[3] P. Dhariwal and A. Q. Nichol, "Diffusion models beat GANs on image synthesis,"
|
||||
in Advances in Neural Information Processing Systems, vol. 34, 2021, pp. 8780-8794.
|
||||
|
||||
3. "classifier-free": classifier-free guidance sampling by conditional DPMs.
|
||||
The input `model` has the following format:
|
||||
``
|
||||
model(x, t_input, cond, **model_kwargs) -> noise | x_start | v | score
|
||||
``
|
||||
And if cond == `unconditional_condition`, the model output is the unconditional DPM output.
|
||||
|
||||
[4] Ho, Jonathan, and Tim Salimans. "Classifier-free diffusion guidance."
|
||||
arXiv preprint arXiv:2207.12598 (2022).
|
||||
|
||||
|
||||
The `t_input` is the time label of the model, which may be discrete-time labels (i.e. 0 to 999)
|
||||
or continuous-time labels (i.e. epsilon to T).
|
||||
|
||||
We wrap the model function to accept only `x` and `t_continuous` as inputs, and outputs the predicted noise:
|
||||
``
|
||||
def model_fn(x, t_continuous) -> noise:
|
||||
t_input = get_model_input_time(t_continuous)
|
||||
return noise_pred(model, x, t_input, **model_kwargs)
|
||||
``
|
||||
where `t_continuous` is the continuous time labels (i.e. epsilon to T). And we use `model_fn` for DPM-Solver.
|
||||
|
||||
===============================================================
|
||||
|
||||
Args:
|
||||
model: A diffusion model with the corresponding format described above.
|
||||
noise_schedule: A noise schedule object, such as NoiseScheduleVP.
|
||||
model_type: A `str`. The parameterization type of the diffusion model.
|
||||
"noise" or "x_start" or "v" or "score".
|
||||
model_kwargs: A `dict`. A dict for the other inputs of the model function.
|
||||
guidance_type: A `str`. The type of the guidance for sampling.
|
||||
"uncond" or "classifier" or "classifier-free".
|
||||
condition: A pytorch tensor. The condition for the guided sampling.
|
||||
Only used for "classifier" or "classifier-free" guidance type.
|
||||
unconditional_condition: A pytorch tensor. The condition for the unconditional sampling.
|
||||
Only used for "classifier-free" guidance type.
|
||||
guidance_scale: A `float`. The scale for the guided sampling.
|
||||
classifier_fn: A classifier function. Only used for the classifier guidance.
|
||||
classifier_kwargs: A `dict`. A dict for the other inputs of the classifier function.
|
||||
Returns:
|
||||
A noise prediction model that accepts the noised data and the continuous time as the inputs.
|
||||
"""
|
||||
|
||||
def get_model_input_time(t_continuous):
|
||||
"""
|
||||
Convert the continuous-time `t_continuous` (in [epsilon, T]) to the model input time.
|
||||
For discrete-time DPMs, we convert `t_continuous` in [1 / N, 1] to `t_input` in [0, 1000 * (N - 1) / N].
|
||||
For continuous-time DPMs, we just use `t_continuous`.
|
||||
"""
|
||||
if noise_schedule.schedule == 'discrete':
|
||||
return (t_continuous - 1. / noise_schedule.total_N) * 1000.
|
||||
else:
|
||||
return t_continuous
|
||||
|
||||
def noise_pred_fn(x, t_continuous, cond=None):
|
||||
if t_continuous.reshape((-1,)).shape[0] == 1:
|
||||
t_continuous = t_continuous.expand((x.shape[0]))
|
||||
t_input = get_model_input_time(t_continuous)
|
||||
if cond is None:
|
||||
output = model(x, t_input, None, **model_kwargs)
|
||||
else:
|
||||
output = model(x, t_input, cond, **model_kwargs)
|
||||
if model_type == "noise":
|
||||
return output
|
||||
elif model_type == "x_start":
|
||||
alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous)
|
||||
dims = x.dim()
|
||||
return (x - expand_dims(alpha_t, dims) * output) / expand_dims(sigma_t, dims)
|
||||
elif model_type == "v":
|
||||
alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous)
|
||||
dims = x.dim()
|
||||
return expand_dims(alpha_t, dims) * output + expand_dims(sigma_t, dims) * x
|
||||
elif model_type == "score":
|
||||
sigma_t = noise_schedule.marginal_std(t_continuous)
|
||||
dims = x.dim()
|
||||
return -expand_dims(sigma_t, dims) * output
|
||||
|
||||
def cond_grad_fn(x, t_input):
|
||||
"""
|
||||
Compute the gradient of the classifier, i.e. nabla_{x} log p_t(cond | x_t).
|
||||
"""
|
||||
with torch.enable_grad():
|
||||
x_in = x.detach().requires_grad_(True)
|
||||
log_prob = classifier_fn(x_in, t_input, condition, **classifier_kwargs)
|
||||
return torch.autograd.grad(log_prob.sum(), x_in)[0]
|
||||
|
||||
def model_fn(x, t_continuous):
|
||||
"""
|
||||
The noise predicition model function that is used for DPM-Solver.
|
||||
"""
|
||||
if t_continuous.reshape((-1,)).shape[0] == 1:
|
||||
t_continuous = t_continuous.expand((x.shape[0]))
|
||||
if guidance_type == "uncond":
|
||||
return noise_pred_fn(x, t_continuous)
|
||||
elif guidance_type == "classifier":
|
||||
assert classifier_fn is not None
|
||||
t_input = get_model_input_time(t_continuous)
|
||||
cond_grad = cond_grad_fn(x, t_input)
|
||||
sigma_t = noise_schedule.marginal_std(t_continuous)
|
||||
noise = noise_pred_fn(x, t_continuous)
|
||||
return noise - guidance_scale * expand_dims(sigma_t, dims=cond_grad.dim()) * cond_grad
|
||||
elif guidance_type == "classifier-free":
|
||||
if guidance_scale == 1. or unconditional_condition is None:
|
||||
return noise_pred_fn(x, t_continuous, cond=condition)
|
||||
else:
|
||||
x_in = torch.cat([x] * 2)
|
||||
t_in = torch.cat([t_continuous] * 2)
|
||||
c_in = torch.cat([unconditional_condition, condition])
|
||||
noise_uncond, noise = noise_pred_fn(x_in, t_in, cond=c_in).chunk(2)
|
||||
return noise_uncond + guidance_scale * (noise - noise_uncond)
|
||||
|
||||
assert model_type in ["noise", "x_start", "v"]
|
||||
assert guidance_type in ["uncond", "classifier", "classifier-free"]
|
||||
return model_fn
|
||||
|
||||
|
||||
class UniPC:
|
||||
def __init__(
|
||||
self,
|
||||
model_fn,
|
||||
noise_schedule,
|
||||
predict_x0=True,
|
||||
thresholding=False,
|
||||
max_val=1.,
|
||||
variant='bh1'
|
||||
):
|
||||
"""Construct a UniPC.
|
||||
|
||||
We support both data_prediction and noise_prediction.
|
||||
"""
|
||||
self.model = model_fn
|
||||
self.noise_schedule = noise_schedule
|
||||
self.variant = variant
|
||||
self.predict_x0 = predict_x0
|
||||
self.thresholding = thresholding
|
||||
self.max_val = max_val
|
||||
|
||||
def dynamic_thresholding_fn(self, x0, t=None):
|
||||
"""
|
||||
The dynamic thresholding method.
|
||||
"""
|
||||
dims = x0.dim()
|
||||
p = self.dynamic_thresholding_ratio
|
||||
s = torch.quantile(torch.abs(x0).reshape((x0.shape[0], -1)), p, dim=1)
|
||||
s = expand_dims(torch.maximum(s, self.thresholding_max_val * torch.ones_like(s).to(s.device)), dims)
|
||||
x0 = torch.clamp(x0, -s, s) / s
|
||||
return x0
|
||||
|
||||
def noise_prediction_fn(self, x, t):
|
||||
"""
|
||||
Return the noise prediction model.
|
||||
"""
|
||||
return self.model(x, t)
|
||||
|
||||
def data_prediction_fn(self, x, t):
|
||||
"""
|
||||
Return the data prediction model (with thresholding).
|
||||
"""
|
||||
noise = self.noise_prediction_fn(x, t)
|
||||
dims = x.dim()
|
||||
alpha_t, sigma_t = self.noise_schedule.marginal_alpha(t), self.noise_schedule.marginal_std(t)
|
||||
x0 = (x - expand_dims(sigma_t, dims) * noise) / expand_dims(alpha_t, dims)
|
||||
if self.thresholding:
|
||||
p = 0.995 # A hyperparameter in the paper of "Imagen" [1].
|
||||
s = torch.quantile(torch.abs(x0).reshape((x0.shape[0], -1)), p, dim=1)
|
||||
s = expand_dims(torch.maximum(s, self.max_val * torch.ones_like(s).to(s.device)), dims)
|
||||
x0 = torch.clamp(x0, -s, s) / s
|
||||
return x0
|
||||
|
||||
def model_fn(self, x, t):
|
||||
"""
|
||||
Convert the model to the noise prediction model or the data prediction model.
|
||||
"""
|
||||
if self.predict_x0:
|
||||
return self.data_prediction_fn(x, t)
|
||||
else:
|
||||
return self.noise_prediction_fn(x, t)
|
||||
|
||||
def get_time_steps(self, skip_type, t_T, t_0, N, device):
|
||||
"""Compute the intermediate time steps for sampling.
|
||||
"""
|
||||
if skip_type == 'logSNR':
|
||||
lambda_T = self.noise_schedule.marginal_lambda(torch.tensor(t_T).to(device))
|
||||
lambda_0 = self.noise_schedule.marginal_lambda(torch.tensor(t_0).to(device))
|
||||
logSNR_steps = torch.linspace(lambda_T.cpu().item(), lambda_0.cpu().item(), N + 1).to(device)
|
||||
return self.noise_schedule.inverse_lambda(logSNR_steps)
|
||||
elif skip_type == 'time_uniform':
|
||||
return torch.linspace(t_T, t_0, N + 1).to(device)
|
||||
elif skip_type == 'time_quadratic':
|
||||
t_order = 2
|
||||
t = torch.linspace(t_T**(1. / t_order), t_0**(1. / t_order), N + 1).pow(t_order).to(device)
|
||||
return t
|
||||
else:
|
||||
raise ValueError("Unsupported skip_type {}, need to be 'logSNR' or 'time_uniform' or 'time_quadratic'".format(skip_type))
|
||||
|
||||
def get_orders_and_timesteps_for_singlestep_solver(self, steps, order, skip_type, t_T, t_0, device):
|
||||
"""
|
||||
Get the order of each step for sampling by the singlestep DPM-Solver.
|
||||
"""
|
||||
if order == 3:
|
||||
K = steps // 3 + 1
|
||||
if steps % 3 == 0:
|
||||
orders = [3,] * (K - 2) + [2, 1]
|
||||
elif steps % 3 == 1:
|
||||
orders = [3,] * (K - 1) + [1]
|
||||
else:
|
||||
orders = [3,] * (K - 1) + [2]
|
||||
elif order == 2:
|
||||
if steps % 2 == 0:
|
||||
K = steps // 2
|
||||
orders = [2,] * K
|
||||
else:
|
||||
K = steps // 2 + 1
|
||||
orders = [2,] * (K - 1) + [1]
|
||||
elif order == 1:
|
||||
K = steps
|
||||
orders = [1,] * steps
|
||||
else:
|
||||
raise ValueError("'order' must be '1' or '2' or '3'.")
|
||||
if skip_type == 'logSNR':
|
||||
# To reproduce the results in DPM-Solver paper
|
||||
timesteps_outer = self.get_time_steps(skip_type, t_T, t_0, K, device)
|
||||
else:
|
||||
timesteps_outer = self.get_time_steps(skip_type, t_T, t_0, steps, device)[torch.cumsum(torch.tensor([0,] + orders), 0).to(device)]
|
||||
return timesteps_outer, orders
|
||||
|
||||
def denoise_to_zero_fn(self, x, s):
|
||||
"""
|
||||
Denoise at the final step, which is equivalent to solve the ODE from lambda_s to infty by first-order discretization.
|
||||
"""
|
||||
return self.data_prediction_fn(x, s)
|
||||
|
||||
def multistep_uni_pc_update(self, x, model_prev_list, t_prev_list, t, order, **kwargs):
|
||||
if len(t.shape) == 0:
|
||||
t = t.view(-1)
|
||||
if 'bh' in self.variant:
|
||||
return self.multistep_uni_pc_bh_update(x, model_prev_list, t_prev_list, t, order, **kwargs)
|
||||
else:
|
||||
assert self.variant == 'vary_coeff'
|
||||
return self.multistep_uni_pc_vary_update(x, model_prev_list, t_prev_list, t, order, **kwargs)
|
||||
|
||||
def multistep_uni_pc_vary_update(self, x, model_prev_list, t_prev_list, t, order, use_corrector=True):
|
||||
print(f'using unified predictor-corrector with order {order} (solver type: vary coeff)')
|
||||
ns = self.noise_schedule
|
||||
assert order <= len(model_prev_list)
|
||||
|
||||
# first compute rks
|
||||
t_prev_0 = t_prev_list[-1]
|
||||
lambda_prev_0 = ns.marginal_lambda(t_prev_0)
|
||||
lambda_t = ns.marginal_lambda(t)
|
||||
model_prev_0 = model_prev_list[-1]
|
||||
sigma_prev_0, sigma_t = ns.marginal_std(t_prev_0), ns.marginal_std(t)
|
||||
log_alpha_t = ns.marginal_log_mean_coeff(t)
|
||||
alpha_t = torch.exp(log_alpha_t)
|
||||
|
||||
h = lambda_t - lambda_prev_0
|
||||
|
||||
rks = []
|
||||
D1s = []
|
||||
for i in range(1, order):
|
||||
t_prev_i = t_prev_list[-(i + 1)]
|
||||
model_prev_i = model_prev_list[-(i + 1)]
|
||||
lambda_prev_i = ns.marginal_lambda(t_prev_i)
|
||||
rk = (lambda_prev_i - lambda_prev_0) / h
|
||||
rks.append(rk)
|
||||
D1s.append((model_prev_i - model_prev_0) / rk)
|
||||
|
||||
rks.append(1.)
|
||||
rks = torch.tensor(rks, device=x.device)
|
||||
|
||||
K = len(rks)
|
||||
# build C matrix
|
||||
C = []
|
||||
|
||||
col = torch.ones_like(rks)
|
||||
for k in range(1, K + 1):
|
||||
C.append(col)
|
||||
col = col * rks / (k + 1)
|
||||
C = torch.stack(C, dim=1)
|
||||
|
||||
if len(D1s) > 0:
|
||||
D1s = torch.stack(D1s, dim=1) # (B, K)
|
||||
C_inv_p = torch.linalg.inv(C[:-1, :-1])
|
||||
A_p = C_inv_p
|
||||
|
||||
if use_corrector:
|
||||
print('using corrector')
|
||||
C_inv = torch.linalg.inv(C)
|
||||
A_c = C_inv
|
||||
|
||||
hh = -h if self.predict_x0 else h
|
||||
h_phi_1 = torch.expm1(hh)
|
||||
h_phi_ks = []
|
||||
factorial_k = 1
|
||||
h_phi_k = h_phi_1
|
||||
for k in range(1, K + 2):
|
||||
h_phi_ks.append(h_phi_k)
|
||||
h_phi_k = h_phi_k / hh - 1 / factorial_k
|
||||
factorial_k *= (k + 1)
|
||||
|
||||
model_t = None
|
||||
if self.predict_x0:
|
||||
x_t_ = (
|
||||
sigma_t / sigma_prev_0 * x
|
||||
- alpha_t * h_phi_1 * model_prev_0
|
||||
)
|
||||
# now predictor
|
||||
x_t = x_t_
|
||||
if len(D1s) > 0:
|
||||
# compute the residuals for predictor
|
||||
for k in range(K - 1):
|
||||
x_t = x_t - alpha_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_p[k])
|
||||
# now corrector
|
||||
if use_corrector:
|
||||
model_t = self.model_fn(x_t, t)
|
||||
D1_t = (model_t - model_prev_0)
|
||||
x_t = x_t_
|
||||
k = 0
|
||||
for k in range(K - 1):
|
||||
x_t = x_t - alpha_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_c[k][:-1])
|
||||
x_t = x_t - alpha_t * h_phi_ks[K] * (D1_t * A_c[k][-1])
|
||||
else:
|
||||
log_alpha_prev_0, log_alpha_t = ns.marginal_log_mean_coeff(t_prev_0), ns.marginal_log_mean_coeff(t)
|
||||
x_t_ = (
|
||||
(torch.exp(log_alpha_t - log_alpha_prev_0)) * x
|
||||
- (sigma_t * h_phi_1) * model_prev_0
|
||||
)
|
||||
# now predictor
|
||||
x_t = x_t_
|
||||
if len(D1s) > 0:
|
||||
# compute the residuals for predictor
|
||||
for k in range(K - 1):
|
||||
x_t = x_t - sigma_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_p[k])
|
||||
# now corrector
|
||||
if use_corrector:
|
||||
model_t = self.model_fn(x_t, t)
|
||||
D1_t = (model_t - model_prev_0)
|
||||
x_t = x_t_
|
||||
k = 0
|
||||
for k in range(K - 1):
|
||||
x_t = x_t - sigma_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_c[k][:-1])
|
||||
x_t = x_t - sigma_t * h_phi_ks[K] * (D1_t * A_c[k][-1])
|
||||
return x_t, model_t
|
||||
|
||||
def multistep_uni_pc_bh_update(self, x, model_prev_list, t_prev_list, t, order, x_t=None, use_corrector=True):
|
||||
print(f'using unified predictor-corrector with order {order} (solver type: B(h))')
|
||||
ns = self.noise_schedule
|
||||
assert order <= len(model_prev_list)
|
||||
dims = x.dim()
|
||||
|
||||
# first compute rks
|
||||
t_prev_0 = t_prev_list[-1]
|
||||
lambda_prev_0 = ns.marginal_lambda(t_prev_0)
|
||||
lambda_t = ns.marginal_lambda(t)
|
||||
model_prev_0 = model_prev_list[-1]
|
||||
sigma_prev_0, sigma_t = ns.marginal_std(t_prev_0), ns.marginal_std(t)
|
||||
log_alpha_prev_0, log_alpha_t = ns.marginal_log_mean_coeff(t_prev_0), ns.marginal_log_mean_coeff(t)
|
||||
alpha_t = torch.exp(log_alpha_t)
|
||||
|
||||
h = lambda_t - lambda_prev_0
|
||||
|
||||
rks = []
|
||||
D1s = []
|
||||
for i in range(1, order):
|
||||
t_prev_i = t_prev_list[-(i + 1)]
|
||||
model_prev_i = model_prev_list[-(i + 1)]
|
||||
lambda_prev_i = ns.marginal_lambda(t_prev_i)
|
||||
rk = ((lambda_prev_i - lambda_prev_0) / h)[0]
|
||||
rks.append(rk)
|
||||
D1s.append((model_prev_i - model_prev_0) / rk)
|
||||
|
||||
rks.append(1.)
|
||||
rks = torch.tensor(rks, device=x.device)
|
||||
|
||||
R = []
|
||||
b = []
|
||||
|
||||
hh = -h[0] if self.predict_x0 else h[0]
|
||||
h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1
|
||||
h_phi_k = h_phi_1 / hh - 1
|
||||
|
||||
factorial_i = 1
|
||||
|
||||
if self.variant == 'bh1':
|
||||
B_h = hh
|
||||
elif self.variant == 'bh2':
|
||||
B_h = torch.expm1(hh)
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
for i in range(1, order + 1):
|
||||
R.append(torch.pow(rks, i - 1))
|
||||
b.append(h_phi_k * factorial_i / B_h)
|
||||
factorial_i *= (i + 1)
|
||||
h_phi_k = h_phi_k / hh - 1 / factorial_i
|
||||
|
||||
R = torch.stack(R)
|
||||
b = torch.tensor(b, device=x.device)
|
||||
|
||||
# now predictor
|
||||
use_predictor = len(D1s) > 0 and x_t is None
|
||||
if len(D1s) > 0:
|
||||
D1s = torch.stack(D1s, dim=1) # (B, K)
|
||||
if x_t is None:
|
||||
# for order 2, we use a simplified version
|
||||
if order == 2:
|
||||
rhos_p = torch.tensor([0.5], device=b.device)
|
||||
else:
|
||||
rhos_p = torch.linalg.solve(R[:-1, :-1], b[:-1])
|
||||
else:
|
||||
D1s = None
|
||||
|
||||
if use_corrector:
|
||||
print('using corrector')
|
||||
# for order 1, we use a simplified version
|
||||
if order == 1:
|
||||
rhos_c = torch.tensor([0.5], device=b.device)
|
||||
else:
|
||||
rhos_c = torch.linalg.solve(R, b)
|
||||
|
||||
model_t = None
|
||||
if self.predict_x0:
|
||||
x_t_ = (
|
||||
expand_dims(sigma_t / sigma_prev_0, dims) * x
|
||||
- expand_dims(alpha_t * h_phi_1, dims)* model_prev_0
|
||||
)
|
||||
|
||||
if x_t is None:
|
||||
if use_predictor:
|
||||
pred_res = torch.einsum('k,bkchw->bchw', rhos_p, D1s)
|
||||
else:
|
||||
pred_res = 0
|
||||
x_t = x_t_ - expand_dims(alpha_t * B_h, dims) * pred_res
|
||||
|
||||
if use_corrector:
|
||||
model_t = self.model_fn(x_t, t)
|
||||
if D1s is not None:
|
||||
corr_res = torch.einsum('k,bkchw->bchw', rhos_c[:-1], D1s)
|
||||
else:
|
||||
corr_res = 0
|
||||
D1_t = (model_t - model_prev_0)
|
||||
x_t = x_t_ - expand_dims(alpha_t * B_h, dims) * (corr_res + rhos_c[-1] * D1_t)
|
||||
else:
|
||||
x_t_ = (
|
||||
expand_dims(torch.exp(log_alpha_t - log_alpha_prev_0), dims) * x
|
||||
- expand_dims(sigma_t * h_phi_1, dims) * model_prev_0
|
||||
)
|
||||
if x_t is None:
|
||||
if use_predictor:
|
||||
pred_res = torch.einsum('k,bkchw->bchw', rhos_p, D1s)
|
||||
else:
|
||||
pred_res = 0
|
||||
x_t = x_t_ - expand_dims(sigma_t * B_h, dims) * pred_res
|
||||
|
||||
if use_corrector:
|
||||
model_t = self.model_fn(x_t, t)
|
||||
if D1s is not None:
|
||||
corr_res = torch.einsum('k,bkchw->bchw', rhos_c[:-1], D1s)
|
||||
else:
|
||||
corr_res = 0
|
||||
D1_t = (model_t - model_prev_0)
|
||||
x_t = x_t_ - expand_dims(sigma_t * B_h, dims) * (corr_res + rhos_c[-1] * D1_t)
|
||||
return x_t, model_t
|
||||
|
||||
|
||||
def sample(self, x, steps=20, t_start=None, t_end=None, order=3, skip_type='time_uniform',
|
||||
method='singlestep', lower_order_final=True, denoise_to_zero=False, solver_type='dpm_solver',
|
||||
atol=0.0078, rtol=0.05, corrector=False,
|
||||
):
|
||||
t_0 = 1. / self.noise_schedule.total_N if t_end is None else t_end
|
||||
t_T = self.noise_schedule.T if t_start is None else t_start
|
||||
device = x.device
|
||||
if method == 'multistep':
|
||||
assert steps >= order
|
||||
timesteps = self.get_time_steps(skip_type=skip_type, t_T=t_T, t_0=t_0, N=steps, device=device)
|
||||
assert timesteps.shape[0] - 1 == steps
|
||||
with torch.no_grad():
|
||||
vec_t = timesteps[0].expand((x.shape[0]))
|
||||
model_prev_list = [self.model_fn(x, vec_t)]
|
||||
t_prev_list = [vec_t]
|
||||
# Init the first `order` values by lower order multistep DPM-Solver.
|
||||
for init_order in range(1, order):
|
||||
vec_t = timesteps[init_order].expand(x.shape[0])
|
||||
x, model_x = self.multistep_uni_pc_update(x, model_prev_list, t_prev_list, vec_t, init_order, use_corrector=True)
|
||||
if model_x is None:
|
||||
model_x = self.model_fn(x, vec_t)
|
||||
model_prev_list.append(model_x)
|
||||
t_prev_list.append(vec_t)
|
||||
for step in range(order, steps + 1):
|
||||
vec_t = timesteps[step].expand(x.shape[0])
|
||||
if lower_order_final:
|
||||
step_order = min(order, steps + 1 - step)
|
||||
else:
|
||||
step_order = order
|
||||
print('this step order:', step_order)
|
||||
if step == steps:
|
||||
print('do not run corrector at the last step')
|
||||
use_corrector = False
|
||||
else:
|
||||
use_corrector = True
|
||||
x, model_x = self.multistep_uni_pc_update(x, model_prev_list, t_prev_list, vec_t, step_order, use_corrector=use_corrector)
|
||||
for i in range(order - 1):
|
||||
t_prev_list[i] = t_prev_list[i + 1]
|
||||
model_prev_list[i] = model_prev_list[i + 1]
|
||||
t_prev_list[-1] = vec_t
|
||||
# We do not need to evaluate the final model value.
|
||||
if step < steps:
|
||||
if model_x is None:
|
||||
model_x = self.model_fn(x, vec_t)
|
||||
model_prev_list[-1] = model_x
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
if denoise_to_zero:
|
||||
x = self.denoise_to_zero_fn(x, torch.ones((x.shape[0],)).to(device) * t_0)
|
||||
return x
|
||||
|
||||
|
||||
#############################################################
|
||||
# other utility functions
|
||||
#############################################################
|
||||
|
||||
def interpolate_fn(x, xp, yp):
|
||||
"""
|
||||
A piecewise linear function y = f(x), using xp and yp as keypoints.
|
||||
We implement f(x) in a differentiable way (i.e. applicable for autograd).
|
||||
The function f(x) is well-defined for all x-axis. (For x beyond the bounds of xp, we use the outmost points of xp to define the linear function.)
|
||||
|
||||
Args:
|
||||
x: PyTorch tensor with shape [N, C], where N is the batch size, C is the number of channels (we use C = 1 for DPM-Solver).
|
||||
xp: PyTorch tensor with shape [C, K], where K is the number of keypoints.
|
||||
yp: PyTorch tensor with shape [C, K].
|
||||
Returns:
|
||||
The function values f(x), with shape [N, C].
|
||||
"""
|
||||
N, K = x.shape[0], xp.shape[1]
|
||||
all_x = torch.cat([x.unsqueeze(2), xp.unsqueeze(0).repeat((N, 1, 1))], dim=2)
|
||||
sorted_all_x, x_indices = torch.sort(all_x, dim=2)
|
||||
x_idx = torch.argmin(x_indices, dim=2)
|
||||
cand_start_idx = x_idx - 1
|
||||
start_idx = torch.where(
|
||||
torch.eq(x_idx, 0),
|
||||
torch.tensor(1, device=x.device),
|
||||
torch.where(
|
||||
torch.eq(x_idx, K), torch.tensor(K - 2, device=x.device), cand_start_idx,
|
||||
),
|
||||
)
|
||||
end_idx = torch.where(torch.eq(start_idx, cand_start_idx), start_idx + 2, start_idx + 1)
|
||||
start_x = torch.gather(sorted_all_x, dim=2, index=start_idx.unsqueeze(2)).squeeze(2)
|
||||
end_x = torch.gather(sorted_all_x, dim=2, index=end_idx.unsqueeze(2)).squeeze(2)
|
||||
start_idx2 = torch.where(
|
||||
torch.eq(x_idx, 0),
|
||||
torch.tensor(0, device=x.device),
|
||||
torch.where(
|
||||
torch.eq(x_idx, K), torch.tensor(K - 2, device=x.device), cand_start_idx,
|
||||
),
|
||||
)
|
||||
y_positions_expanded = yp.unsqueeze(0).expand(N, -1, -1)
|
||||
start_y = torch.gather(y_positions_expanded, dim=2, index=start_idx2.unsqueeze(2)).squeeze(2)
|
||||
end_y = torch.gather(y_positions_expanded, dim=2, index=(start_idx2 + 1).unsqueeze(2)).squeeze(2)
|
||||
cand = start_y + (x - start_x) * (end_y - start_y) / (end_x - start_x)
|
||||
return cand
|
||||
|
||||
|
||||
def expand_dims(v, dims):
|
||||
"""
|
||||
Expand the tensor `v` to the dim `dims`.
|
||||
|
||||
Args:
|
||||
`v`: a PyTorch tensor with shape [N].
|
||||
`dim`: a `int`.
|
||||
Returns:
|
||||
a PyTorch tensor with shape [N, 1, 1, ..., 1] and the total dimension is `dims`.
|
||||
"""
|
||||
return v[(...,) + (None,)*(dims - 1)]
|
||||
@@ -0,0 +1,158 @@
|
||||
import math
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import repeat
|
||||
|
||||
|
||||
def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False, dtype=None):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param timesteps: a 1-D Tensor of N indices, one per batch element.
|
||||
These may be fractional.
|
||||
:param dim: the dimension of the output.
|
||||
:param max_period: controls the minimum frequency of the embeddings.
|
||||
:return: an [N x dim] Tensor of positional embeddings.
|
||||
"""
|
||||
if not repeat_only:
|
||||
half = dim // 2
|
||||
freqs = torch.exp(
|
||||
-math.log(max_period) * torch.arange(start=0, end=half, dtype=dtype) / half
|
||||
).to(device=timesteps.device)
|
||||
args = timesteps[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
else:
|
||||
embedding = repeat(timesteps, 'b -> b d', d=dim)
|
||||
return embedding.to(dtype)
|
||||
|
||||
|
||||
def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3):
|
||||
if schedule == "linear":
|
||||
betas = (
|
||||
torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=torch.float64) ** 2
|
||||
)
|
||||
|
||||
elif schedule == "cosine":
|
||||
timesteps = (
|
||||
torch.arange(n_timestep + 1, dtype=torch.float64) / n_timestep + cosine_s
|
||||
)
|
||||
alphas = timesteps / (1 + cosine_s) * np.pi / 2
|
||||
alphas = torch.cos(alphas).pow(2)
|
||||
alphas = alphas / alphas[0]
|
||||
betas = 1 - alphas[1:] / alphas[:-1]
|
||||
betas = np.clip(betas, a_min=0, a_max=0.999)
|
||||
|
||||
elif schedule == "sqrt_linear":
|
||||
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64)
|
||||
elif schedule == "sqrt":
|
||||
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64) ** 0.5
|
||||
else:
|
||||
raise ValueError(f"schedule '{schedule}' unknown.")
|
||||
return betas.numpy()
|
||||
|
||||
|
||||
def make_ddim_timesteps(ddim_discr_method, num_ddim_timesteps, num_ddpm_timesteps, verbose=True):
|
||||
if ddim_discr_method == 'uniform':
|
||||
c = num_ddpm_timesteps // num_ddim_timesteps
|
||||
ddim_timesteps = np.asarray(list(range(0, num_ddpm_timesteps, c)))
|
||||
steps_out = ddim_timesteps + 1
|
||||
elif ddim_discr_method == 'uniform_trailing':
|
||||
c = num_ddpm_timesteps / num_ddim_timesteps
|
||||
ddim_timesteps = np.flip(np.round(np.arange(num_ddpm_timesteps, 0, -c))).astype(np.int64)
|
||||
steps_out = ddim_timesteps - 1
|
||||
elif ddim_discr_method == 'quad':
|
||||
ddim_timesteps = ((np.linspace(0, np.sqrt(num_ddpm_timesteps * .8), num_ddim_timesteps)) ** 2).astype(int)
|
||||
steps_out = ddim_timesteps + 1
|
||||
else:
|
||||
raise NotImplementedError(f'There is no ddim discretization method called "{ddim_discr_method}"')
|
||||
|
||||
# assert ddim_timesteps.shape[0] == num_ddim_timesteps
|
||||
# add one to get the final alpha values right (the ones from first scale to data during sampling)
|
||||
# steps_out = ddim_timesteps + 1
|
||||
if verbose:
|
||||
print(f'Selected timesteps for ddim sampler: {steps_out}')
|
||||
return steps_out
|
||||
|
||||
|
||||
def make_ddim_sampling_parameters(alphacums, ddim_timesteps, eta, verbose=True):
|
||||
# select alphas for computing the variance schedule
|
||||
# print(f'ddim_timesteps={ddim_timesteps}, len_alphacums={len(alphacums)}')
|
||||
alphas = alphacums[ddim_timesteps]
|
||||
alphas_prev = np.asarray([alphacums[0]] + alphacums[ddim_timesteps[:-1]].tolist())
|
||||
|
||||
# according the the formula provided in https://arxiv.org/abs/2010.02502
|
||||
sigmas = eta * np.sqrt((1 - alphas_prev) / (1 - alphas) * (1 - alphas / alphas_prev))
|
||||
if verbose:
|
||||
print(f'Selected alphas for ddim sampler: a_t: {alphas}; a_(t-1): {alphas_prev}')
|
||||
print(f'For the chosen value of eta, which is {eta}, '
|
||||
f'this results in the following sigma_t schedule for ddim sampler {sigmas}')
|
||||
return sigmas, alphas, alphas_prev
|
||||
|
||||
|
||||
def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
|
||||
"""
|
||||
Create a beta schedule that discretizes the given alpha_t_bar function,
|
||||
which defines the cumulative product of (1-beta) over time from t = [0,1].
|
||||
:param num_diffusion_timesteps: the number of betas to produce.
|
||||
:param alpha_bar: a lambda that takes an argument t from 0 to 1 and
|
||||
produces the cumulative product of (1-beta) up to that
|
||||
part of the diffusion process.
|
||||
:param max_beta: the maximum beta to use; use values lower than 1 to
|
||||
prevent singularities.
|
||||
"""
|
||||
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 np.array(betas)
|
||||
|
||||
def rescale_zero_terminal_snr(betas):
|
||||
"""
|
||||
Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1)
|
||||
|
||||
Args:
|
||||
betas (`numpy.ndarray`):
|
||||
the betas that the scheduler is being initialized with.
|
||||
|
||||
Returns:
|
||||
`numpy.ndarray`: rescaled betas with zero terminal SNR
|
||||
"""
|
||||
# Convert betas to alphas_bar_sqrt
|
||||
alphas = 1.0 - betas
|
||||
alphas_cumprod = np.cumprod(alphas, axis=0)
|
||||
alphas_bar_sqrt = np.sqrt(alphas_cumprod)
|
||||
|
||||
# Store old values.
|
||||
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].copy()
|
||||
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].copy()
|
||||
|
||||
# Shift so the last timestep is zero.
|
||||
alphas_bar_sqrt -= alphas_bar_sqrt_T
|
||||
|
||||
# Scale so the first timestep is back to the old value.
|
||||
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
|
||||
|
||||
# Convert alphas_bar_sqrt to betas
|
||||
alphas_bar = alphas_bar_sqrt**2 # Revert sqrt
|
||||
alphas = alphas_bar[1:] / alphas_bar[:-1] # Revert cumprod
|
||||
alphas = np.concatenate([alphas_bar[0:1], alphas])
|
||||
betas = 1 - alphas
|
||||
|
||||
return betas
|
||||
|
||||
|
||||
def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0):
|
||||
"""
|
||||
Rescale `noise_cfg` according to `guidance_rescale`. Based on findings of [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). See Section 3.4
|
||||
"""
|
||||
std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True)
|
||||
std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)
|
||||
# rescale the results from guidance (fixes overexposure)
|
||||
noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
|
||||
# mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images
|
||||
noise_cfg = guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg
|
||||
return noise_cfg
|
||||
@@ -0,0 +1,809 @@
|
||||
import torch
|
||||
from torch import nn, einsum
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange, repeat
|
||||
from functools import partial
|
||||
from ..common import (
|
||||
checkpoint,
|
||||
exists,
|
||||
default,
|
||||
)
|
||||
from ..basics import zero_module
|
||||
import comfy.ops
|
||||
ops = comfy.ops.disable_weight_init
|
||||
from comfy import model_management
|
||||
from comfy.ldm.modules.attention import optimized_attention, optimized_attention_masked
|
||||
|
||||
if model_management.xformers_enabled():
|
||||
import xformers
|
||||
import xformers.ops
|
||||
XFORMERS_IS_AVAILBLE = True
|
||||
else:
|
||||
XFORMERS_IS_AVAILBLE = False
|
||||
|
||||
class RelativePosition(nn.Module):
|
||||
""" https://github.com/evelinehong/Transformer_Relative_Position_PyTorch/blob/master/relative_position.py """
|
||||
|
||||
def __init__(self, num_units, max_relative_position):
|
||||
super().__init__()
|
||||
self.num_units = num_units
|
||||
self.max_relative_position = max_relative_position
|
||||
self.embeddings_table = nn.Parameter(torch.Tensor(max_relative_position * 2 + 1, num_units))
|
||||
nn.init.xavier_uniform_(self.embeddings_table)
|
||||
|
||||
def forward(self, length_q, length_k):
|
||||
device = self.embeddings_table.device
|
||||
range_vec_q = torch.arange(length_q, device=device)
|
||||
range_vec_k = torch.arange(length_k, device=device)
|
||||
distance_mat = range_vec_k[None, :] - range_vec_q[:, None]
|
||||
distance_mat_clipped = torch.clamp(distance_mat, -self.max_relative_position, self.max_relative_position)
|
||||
final_mat = distance_mat_clipped + self.max_relative_position
|
||||
final_mat = final_mat.long()
|
||||
embeddings = self.embeddings_table[final_mat]
|
||||
return embeddings
|
||||
|
||||
|
||||
# TODO Add native Comfy optimized attention.
|
||||
class CrossAttention(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
query_dim,
|
||||
context_dim=None,
|
||||
heads=8,
|
||||
dim_head=64,
|
||||
dropout=0.,
|
||||
relative_position=False,
|
||||
temporal_length=None,
|
||||
video_length=None,
|
||||
image_cross_attention=False,
|
||||
image_cross_attention_scale=1.0,
|
||||
image_cross_attention_scale_learnable=False,
|
||||
text_context_len=77,
|
||||
device=None,
|
||||
dtype=None,
|
||||
operations=ops
|
||||
):
|
||||
super().__init__()
|
||||
inner_dim = dim_head * heads
|
||||
context_dim = default(context_dim, query_dim)
|
||||
self.scale = dim_head**-0.5
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
self.to_q = operations.Linear(query_dim, inner_dim, bias=False, device=device, dtype=dtype)
|
||||
self.to_k = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype)
|
||||
self.to_v = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
operations.Linear(inner_dim, query_dim, device=device, dtype=dtype),
|
||||
nn.Dropout(dropout)
|
||||
)
|
||||
|
||||
self.relative_position = relative_position
|
||||
if self.relative_position:
|
||||
assert(temporal_length is not None)
|
||||
self.relative_position_k = RelativePosition(num_units=dim_head, max_relative_position=temporal_length)
|
||||
self.relative_position_v = RelativePosition(num_units=dim_head, max_relative_position=temporal_length)
|
||||
else:
|
||||
## only used for spatial attention, while NOT for temporal attention
|
||||
if XFORMERS_IS_AVAILBLE and temporal_length is None:
|
||||
self.forward = self.efficient_forward
|
||||
else:
|
||||
self.forward = self.comfy_efficient_forward
|
||||
|
||||
self.video_length = video_length
|
||||
self.image_cross_attention = image_cross_attention
|
||||
self.image_cross_attention_scale = image_cross_attention_scale
|
||||
self.text_context_len = text_context_len
|
||||
self.image_cross_attention_scale_learnable = image_cross_attention_scale_learnable
|
||||
if self.image_cross_attention:
|
||||
self.to_k_ip = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype)
|
||||
self.to_v_ip = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype)
|
||||
if image_cross_attention_scale_learnable:
|
||||
self.register_parameter('alpha', nn.Parameter(torch.tensor(0.)) )
|
||||
|
||||
def comfy_efficient_forward(self, x, context=None, mask=None, *args, **kwargs):
|
||||
spatial_self_attn = (context is None)
|
||||
k_ip, v_ip, out_ip = None, None, None
|
||||
|
||||
h = self.heads
|
||||
q = self.to_q(x)
|
||||
context = default(context, x)
|
||||
|
||||
if self.image_cross_attention and not spatial_self_attn:
|
||||
context, context_image = context[:,:self.text_context_len,:], context[:,self.text_context_len:,:]
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
k_ip = self.to_k_ip(context_image)
|
||||
v_ip = self.to_v_ip(context_image)
|
||||
else:
|
||||
if not spatial_self_attn:
|
||||
context = context[:,:self.text_context_len,:]
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
|
||||
out = optimized_attention(q, k, v, h)
|
||||
|
||||
if exists(mask):
|
||||
## feasible for causal attention mask only
|
||||
out = optimized_attention_masked(q, k, v, h)
|
||||
|
||||
## for image cross-attention
|
||||
if k_ip is not None:
|
||||
q = rearrange(q, 'b n (h d) -> (b h) n d', h=h)
|
||||
k_ip, v_ip = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (k_ip, v_ip))
|
||||
sim_ip = torch.einsum('b i d, b j d -> b i j', q, k_ip) * self.scale
|
||||
del k_ip
|
||||
sim_ip = sim_ip.softmax(dim=-1)
|
||||
out_ip = torch.einsum('b i j, b j d -> b i d', sim_ip, v_ip)
|
||||
out_ip = rearrange(out_ip, '(b h) n d -> b n (h d)', h=h)
|
||||
|
||||
if out_ip is not None:
|
||||
if self.image_cross_attention_scale_learnable:
|
||||
out = out + self.image_cross_attention_scale * out_ip * (torch.tanh(self.alpha)+1)
|
||||
else:
|
||||
out = out + self.image_cross_attention_scale * out_ip
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
def forward(self, x, context=None, mask=None):
|
||||
spatial_self_attn = (context is None)
|
||||
k_ip, v_ip, out_ip = None, None, None
|
||||
|
||||
h = self.heads
|
||||
q = self.to_q(x)
|
||||
context = default(context, x)
|
||||
|
||||
if self.image_cross_attention and not spatial_self_attn:
|
||||
context, context_image = context[:,:self.text_context_len,:], context[:,self.text_context_len:,:]
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
k_ip = self.to_k_ip(context_image)
|
||||
v_ip = self.to_v_ip(context_image)
|
||||
else:
|
||||
|
||||
# Assumed Spatial Attention (b c h w)
|
||||
if not spatial_self_attn:
|
||||
context = context[:,:self.text_context_len,:]
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
|
||||
|
||||
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v))
|
||||
|
||||
sim = torch.einsum('b i d, b j d -> b i j', q, k) * self.scale
|
||||
if self.relative_position:
|
||||
len_q, len_k, len_v = q.shape[1], k.shape[1], v.shape[1]
|
||||
k2 = self.relative_position_k(len_q, len_k)
|
||||
sim2 = einsum('b t d, t s d -> b t s', q, k2) * self.scale # TODO check
|
||||
sim += sim2
|
||||
del k
|
||||
|
||||
if exists(mask):
|
||||
## feasible for causal attention mask only
|
||||
max_neg_value = -torch.finfo(sim.dtype).max
|
||||
mask = repeat(mask, 'b i j -> (b h) i j', h=h)
|
||||
sim.masked_fill_(~(mask>0.5), max_neg_value)
|
||||
|
||||
# attention, what we cannot get enough of
|
||||
sim = sim.softmax(dim=-1)
|
||||
|
||||
out = torch.einsum('b i j, b j d -> b i d', sim, v)
|
||||
if self.relative_position:
|
||||
v2 = self.relative_position_v(len_q, len_v)
|
||||
out2 = einsum('b t s, t s d -> b t d', sim, v2) # TODO check
|
||||
out += out2
|
||||
out = rearrange(out, '(b h) n d -> b n (h d)', h=h)
|
||||
|
||||
|
||||
## for image cross-attention
|
||||
if k_ip is not None:
|
||||
k_ip, v_ip = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (k_ip, v_ip))
|
||||
sim_ip = torch.einsum('b i d, b j d -> b i j', q, k_ip) * self.scale
|
||||
del k_ip
|
||||
sim_ip = sim_ip.softmax(dim=-1)
|
||||
out_ip = torch.einsum('b i j, b j d -> b i d', sim_ip, v_ip)
|
||||
out_ip = rearrange(out_ip, '(b h) n d -> b n (h d)', h=h)
|
||||
|
||||
|
||||
if out_ip is not None:
|
||||
if self.image_cross_attention_scale_learnable:
|
||||
out = out + self.image_cross_attention_scale * out_ip * (torch.tanh(self.alpha)+1)
|
||||
else:
|
||||
out = out + self.image_cross_attention_scale * out_ip
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
def efficient_forward(self, x, context=None, mask=None):
|
||||
spatial_self_attn = (context is None)
|
||||
k_ip, v_ip, out_ip = None, None, None
|
||||
|
||||
q = self.to_q(x)
|
||||
context = default(context, x)
|
||||
|
||||
if self.image_cross_attention and not spatial_self_attn:
|
||||
context, context_image = context[:,:self.text_context_len,:], context[:,self.text_context_len:,:]
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
k_ip = self.to_k_ip(context_image)
|
||||
v_ip = self.to_v_ip(context_image)
|
||||
else:
|
||||
if not spatial_self_attn:
|
||||
context = context[:,:self.text_context_len,:]
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
|
||||
b, _, _ = q.shape
|
||||
q, k, v = map(
|
||||
lambda t: t.unsqueeze(3)
|
||||
.reshape(b, t.shape[1], self.heads, self.dim_head)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(b * self.heads, t.shape[1], self.dim_head)
|
||||
.contiguous(),
|
||||
(q, k, v),
|
||||
)
|
||||
# actually compute the attention, what we cannot get enough of
|
||||
out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=None, op=None)
|
||||
|
||||
## for image cross-attention
|
||||
if k_ip is not None:
|
||||
k_ip, v_ip = map(
|
||||
lambda t: t.unsqueeze(3)
|
||||
.reshape(b, t.shape[1], self.heads, self.dim_head)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(b * self.heads, t.shape[1], self.dim_head)
|
||||
.contiguous(),
|
||||
(k_ip, v_ip),
|
||||
)
|
||||
out_ip = xformers.ops.memory_efficient_attention(q, k_ip, v_ip, attn_bias=None, op=None)
|
||||
out_ip = (
|
||||
out_ip.unsqueeze(0)
|
||||
.reshape(b, self.heads, out.shape[1], self.dim_head)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(b, out.shape[1], self.heads * self.dim_head)
|
||||
)
|
||||
|
||||
if exists(mask):
|
||||
raise NotImplementedError
|
||||
out = (
|
||||
out.unsqueeze(0)
|
||||
.reshape(b, self.heads, out.shape[1], self.dim_head)
|
||||
.permute(0, 2, 1, 3)
|
||||
.reshape(b, out.shape[1], self.heads * self.dim_head)
|
||||
)
|
||||
if out_ip is not None:
|
||||
if self.image_cross_attention_scale_learnable:
|
||||
out = out + self.image_cross_attention_scale * out_ip * (torch.tanh(self.alpha)+1)
|
||||
else:
|
||||
out = out + self.image_cross_attention_scale * out_ip
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class BasicTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
n_heads,
|
||||
d_head,
|
||||
dropout=0.,
|
||||
context_dim=None,
|
||||
gated_ff=True,
|
||||
checkpoint=True,
|
||||
disable_self_attn=False,
|
||||
attention_cls=None,
|
||||
video_length=None,
|
||||
inner_dim=None,
|
||||
image_cross_attention=False,
|
||||
image_cross_attention_scale=1.0,
|
||||
image_cross_attention_scale_learnable=False,
|
||||
switch_temporal_ca_to_sa=False,
|
||||
text_context_len=77,
|
||||
ff_in=None,
|
||||
device=None,
|
||||
dtype=None,
|
||||
operations=ops
|
||||
):
|
||||
super().__init__()
|
||||
attn_cls = CrossAttention if attention_cls is None else attention_cls
|
||||
|
||||
self.ff_in = ff_in or inner_dim is not None
|
||||
if self.ff_in:
|
||||
self.norm_in = operations.LayerNorm(dim, dtype=dtype, device=device)
|
||||
self.ff_in = FeedForward(
|
||||
dim,
|
||||
dim_out=inner_dim,
|
||||
dropout=dropout,
|
||||
glu=gated_ff,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations
|
||||
)
|
||||
if inner_dim is None:
|
||||
inner_dim = dim
|
||||
|
||||
self.is_res = inner_dim == dim
|
||||
self.disable_self_attn = disable_self_attn
|
||||
self.attn1 = attn_cls(query_dim=dim, heads=n_heads, dim_head=d_head, dropout=dropout,
|
||||
context_dim=None, device=device, dtype=dtype if self.disable_self_attn else None)
|
||||
self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff, device=device, dtype=dtype)
|
||||
self.attn2 = attn_cls(
|
||||
query_dim=dim,
|
||||
context_dim=context_dim,
|
||||
heads=n_heads,
|
||||
dim_head=d_head,
|
||||
dropout=dropout,
|
||||
video_length=video_length,
|
||||
image_cross_attention=image_cross_attention,
|
||||
image_cross_attention_scale=image_cross_attention_scale,
|
||||
image_cross_attention_scale_learnable=image_cross_attention_scale_learnable,
|
||||
text_context_len=text_context_len,
|
||||
device=device,
|
||||
dtype=dtype
|
||||
)
|
||||
self.image_cross_attention = image_cross_attention
|
||||
|
||||
self.norm1 = operations.LayerNorm(dim, device=device, dtype=dtype)
|
||||
self.norm2 = operations.LayerNorm(dim, device=device, dtype=dtype)
|
||||
self.norm3 = operations.LayerNorm(dim, device=device, dtype=dtype)
|
||||
|
||||
self.n_heads = n_heads
|
||||
self.d_head = d_head
|
||||
self.checkpoint = checkpoint
|
||||
self.switch_temporal_ca_to_sa = switch_temporal_ca_to_sa
|
||||
|
||||
def forward(self, x, context=None, mask=None, **kwargs):
|
||||
## implementation tricks: because checkpointing doesn't support non-tensor (e.g. None or scalar) arguments
|
||||
input_tuple = (x,) ## should not be (x), otherwise *input_tuple will decouple x into multiple arguments
|
||||
if context is not None:
|
||||
input_tuple = (x, context)
|
||||
if mask is not None:
|
||||
forward_mask = partial(self._forward, mask=mask)
|
||||
return checkpoint(forward_mask, (x,), self.parameters(), self.checkpoint)
|
||||
return checkpoint(self._forward, input_tuple, self.parameters(), self.checkpoint)
|
||||
|
||||
|
||||
def _forward(self, x, context=None, mask=None, transformer_options={}):
|
||||
extra_options = {}
|
||||
block = transformer_options.get("block", None)
|
||||
block_index = transformer_options.get("block_index", 0)
|
||||
transformer_patches = {}
|
||||
transformer_patches_replace = {}
|
||||
|
||||
for k in transformer_options:
|
||||
if k == "patches":
|
||||
transformer_patches = transformer_options[k]
|
||||
elif k == "patches_replace":
|
||||
transformer_patches_replace = transformer_options[k]
|
||||
else:
|
||||
extra_options[k] = transformer_options[k]
|
||||
|
||||
extra_options["n_heads"] = self.n_heads
|
||||
extra_options["dim_head"] = self.d_head
|
||||
|
||||
if self.ff_in:
|
||||
x_skip = x
|
||||
x = self.ff_in(self.norm_in(x))
|
||||
if self.is_res:
|
||||
x += x_skip
|
||||
|
||||
n = self.norm1(x)
|
||||
if self.disable_self_attn:
|
||||
context_attn1 = context
|
||||
else:
|
||||
context_attn1 = None
|
||||
value_attn1 = None
|
||||
|
||||
if "attn1_patch" in transformer_patches:
|
||||
patch = transformer_patches["attn1_patch"]
|
||||
if context_attn1 is None:
|
||||
context_attn1 = n
|
||||
value_attn1 = context_attn1
|
||||
for p in patch:
|
||||
n, context_attn1, value_attn1 = p(n, context_attn1, value_attn1, extra_options)
|
||||
|
||||
if block is not None:
|
||||
transformer_block = (block[0], block[1], block_index)
|
||||
else:
|
||||
transformer_block = None
|
||||
attn1_replace_patch = transformer_patches_replace.get("attn1", {})
|
||||
block_attn1 = transformer_block
|
||||
if block_attn1 not in attn1_replace_patch:
|
||||
block_attn1 = block
|
||||
|
||||
if block_attn1 in attn1_replace_patch:
|
||||
if context_attn1 is None:
|
||||
context_attn1 = n
|
||||
value_attn1 = n
|
||||
n = self.attn1.to_q(n)
|
||||
context_attn1 = self.attn1.to_k(context_attn1)
|
||||
value_attn1 = self.attn1.to_v(value_attn1)
|
||||
n = attn1_replace_patch[block_attn1](n, context_attn1, value_attn1, extra_options)
|
||||
n = self.attn1.to_out(n)
|
||||
else:
|
||||
n = self.attn1(n, context=context_attn1, value=value_attn1)
|
||||
|
||||
if "attn1_output_patch" in transformer_patches:
|
||||
patch = transformer_patches["attn1_output_patch"]
|
||||
for p in patch:
|
||||
n = p(n, extra_options)
|
||||
|
||||
x += n
|
||||
if "middle_patch" in transformer_patches:
|
||||
patch = transformer_patches["middle_patch"]
|
||||
for p in patch:
|
||||
x = p(x, extra_options)
|
||||
|
||||
if self.attn2 is not None:
|
||||
n = self.norm2(x)
|
||||
if self.switch_temporal_ca_to_sa:
|
||||
context_attn2 = n
|
||||
else:
|
||||
context_attn2 = context
|
||||
value_attn2 = None
|
||||
if "attn2_patch" in transformer_patches:
|
||||
patch = transformer_patches["attn2_patch"]
|
||||
value_attn2 = context_attn2
|
||||
for p in patch:
|
||||
n, context_attn2, value_attn2 = p(n, context_attn2, value_attn2, extra_options)
|
||||
|
||||
attn2_replace_patch = transformer_patches_replace.get("attn2", {})
|
||||
block_attn2 = transformer_block
|
||||
if block_attn2 not in attn2_replace_patch:
|
||||
block_attn2 = block
|
||||
|
||||
if block_attn2 in attn2_replace_patch:
|
||||
if value_attn2 is None:
|
||||
value_attn2 = context_attn2
|
||||
n = self.attn2.to_q(n)
|
||||
context_attn2 = self.attn2.to_k(context_attn2)
|
||||
value_attn2 = self.attn2.to_v(value_attn2)
|
||||
n = attn2_replace_patch[block_attn2](n, context_attn2, value_attn2, extra_options)
|
||||
n = self.attn2.to_out(n)
|
||||
else:
|
||||
n = self.attn2(n, context=context_attn2, value=value_attn2)
|
||||
|
||||
if "attn2_output_patch" in transformer_patches:
|
||||
patch = transformer_patches["attn2_output_patch"]
|
||||
for p in patch:
|
||||
n = p(n, extra_options)
|
||||
|
||||
x += n
|
||||
if self.is_res:
|
||||
x_skip = x
|
||||
x = self.ff(self.norm3(x))
|
||||
if self.is_res:
|
||||
x += x_skip
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class SpatialTransformer(nn.Module):
|
||||
"""
|
||||
Transformer block for image-like data in spatial axis.
|
||||
First, project the input (aka embedding)
|
||||
and reshape to b, t, d.
|
||||
Then apply standard transformer action.
|
||||
Finally, reshape to image
|
||||
NEW: use_linear for more efficiency instead of the 1x1 convs
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
n_heads,
|
||||
d_head,
|
||||
depth=1,
|
||||
dropout=0.,
|
||||
context_dim=None,
|
||||
use_checkpoint=True,
|
||||
disable_self_attn=False,
|
||||
use_linear=False,
|
||||
video_length=None,
|
||||
image_cross_attention=False,
|
||||
image_cross_attention_scale_learnable=False,
|
||||
device=None,
|
||||
dtype=None,
|
||||
operations=ops
|
||||
):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
inner_dim = n_heads * d_head
|
||||
self.norm = operations.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True, device=device, dtype=dtype)
|
||||
if not use_linear:
|
||||
self.proj_in = opeations.Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0, device=device, dtype=dtype)
|
||||
else:
|
||||
self.proj_in = operations.Linear(in_channels, inner_dim, device=device, dtype=dtype)
|
||||
|
||||
attention_cls = None
|
||||
self.transformer_blocks = nn.ModuleList([
|
||||
BasicTransformerBlock(
|
||||
inner_dim,
|
||||
n_heads,
|
||||
d_head,
|
||||
dropout=dropout,
|
||||
context_dim=context_dim,
|
||||
disable_self_attn=disable_self_attn,
|
||||
checkpoint=use_checkpoint,
|
||||
attention_cls=attention_cls,
|
||||
video_length=video_length,
|
||||
image_cross_attention=image_cross_attention,
|
||||
image_cross_attention_scale_learnable=image_cross_attention_scale_learnable,
|
||||
device=device,
|
||||
dtype=dtype
|
||||
) for d in range(depth)
|
||||
])
|
||||
if not use_linear:
|
||||
self.proj_out = zero_module(operations.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0, device=device, dtype=dtype))
|
||||
else:
|
||||
self.proj_out = zero_module(operations.Linear(inner_dim, in_channels, device=device, dtype=dtype))
|
||||
self.use_linear = use_linear
|
||||
|
||||
def forward(self, x, context=None, transformer_options={}, **kwargs):
|
||||
b, c, h, w = x.shape
|
||||
x_in = x
|
||||
x = self.norm(x)
|
||||
if not self.use_linear:
|
||||
x = self.proj_in(x)
|
||||
x = rearrange(x, 'b c h w -> b (h w) c').contiguous()
|
||||
if self.use_linear:
|
||||
x = self.proj_in(x)
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
transformer_options['block_index'] = i
|
||||
x = block(x, context=context, **kwargs)
|
||||
if self.use_linear:
|
||||
x = self.proj_out(x)
|
||||
x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous()
|
||||
if not self.use_linear:
|
||||
x = self.proj_out(x)
|
||||
return x + x_in
|
||||
|
||||
|
||||
class TemporalTransformer(nn.Module):
|
||||
"""
|
||||
Transformer block for image-like data in temporal axis.
|
||||
First, reshape to b, t, d.
|
||||
Then apply standard transformer action.
|
||||
Finally, reshape to image
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
n_heads,
|
||||
d_head,
|
||||
depth=1,
|
||||
dropout=0.,
|
||||
context_dim=None,
|
||||
use_checkpoint=True,
|
||||
use_linear=False,
|
||||
only_self_att=True,
|
||||
causal_attention=False,
|
||||
causal_block_size=1,
|
||||
relative_position=False,
|
||||
temporal_length=None,
|
||||
device=None,
|
||||
dtype=None,
|
||||
operations=ops
|
||||
):
|
||||
super().__init__()
|
||||
self.only_self_att = only_self_att
|
||||
self.relative_position = relative_position
|
||||
self.causal_attention = causal_attention
|
||||
self.causal_block_size = causal_block_size
|
||||
|
||||
if only_self_att:
|
||||
context_dim = None
|
||||
|
||||
self.in_channels = in_channels
|
||||
inner_dim = n_heads * d_head
|
||||
self.norm = operations.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True, device=device, dtype=dtype)
|
||||
self.proj_in = nn.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0).to(device, dtype)
|
||||
if not use_linear:
|
||||
self.proj_in = nn.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0).to(device, dtype)
|
||||
else:
|
||||
self.proj_in = operations.Linear(in_channels, inner_dim, device=device, dtype=dtype)
|
||||
|
||||
if relative_position:
|
||||
assert(temporal_length is not None)
|
||||
attention_cls = partial(CrossAttention, relative_position=True, temporal_length=temporal_length, device=device, dtype=dtype)
|
||||
else:
|
||||
attention_cls = partial(CrossAttention, temporal_length=temporal_length, device=device, dtype=dtype)
|
||||
if self.causal_attention:
|
||||
assert(temporal_length is not None)
|
||||
self.mask = torch.tril(torch.ones([1, temporal_length, temporal_length]))
|
||||
|
||||
if self.only_self_att:
|
||||
context_dim = None
|
||||
self.transformer_blocks = nn.ModuleList([
|
||||
BasicTransformerBlock(
|
||||
inner_dim,
|
||||
n_heads,
|
||||
d_head,
|
||||
dropout=dropout,
|
||||
context_dim=context_dim,
|
||||
attention_cls=attention_cls,
|
||||
checkpoint=use_checkpoint,
|
||||
device=device,
|
||||
dtype=dtype
|
||||
) for d in range(depth)
|
||||
])
|
||||
if not use_linear:
|
||||
self.proj_out = zero_module(nn.Conv1d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0).to(device, dtype))
|
||||
else:
|
||||
self.proj_out = zero_module(operations.Linear(inner_dim, in_channels, device=device, dtype=dtype))
|
||||
self.use_linear = use_linear
|
||||
|
||||
def forward(self, x, context=None):
|
||||
b, c, t, h, w = x.shape
|
||||
x_in = x
|
||||
x = self.norm(x)
|
||||
x = rearrange(x, 'b c t h w -> (b h w) c t').contiguous()
|
||||
if not self.use_linear:
|
||||
x = self.proj_in(x)
|
||||
x = rearrange(x, 'bhw c t -> bhw t c').contiguous()
|
||||
if self.use_linear:
|
||||
x = self.proj_in(x)
|
||||
|
||||
temp_mask = None
|
||||
if self.causal_attention:
|
||||
# slice the from mask map
|
||||
temp_mask = self.mask[:,:t,:t].to(x.device)
|
||||
|
||||
if temp_mask is not None:
|
||||
mask = temp_mask.to(x.device)
|
||||
mask = repeat(mask, 'l i j -> (l bhw) i j', bhw=b*h*w)
|
||||
else:
|
||||
mask = None
|
||||
|
||||
if self.only_self_att:
|
||||
## note: if no context is given, cross-attention defaults to self-attention
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
x = block(x, mask=mask)
|
||||
x = rearrange(x, '(b hw) t c -> b hw t c', b=b).contiguous()
|
||||
else:
|
||||
x = rearrange(x, '(b hw) t c -> b hw t c', b=b).contiguous()
|
||||
context = rearrange(context, '(b t) l con -> b t l con', t=t).contiguous()
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
# calculate each batch one by one (since number in shape could not greater then 65,535 for some package)
|
||||
for j in range(b):
|
||||
context_j = repeat(
|
||||
context[j],
|
||||
't l con -> (t r) l con', r=(h * w) // t, t=t).contiguous()
|
||||
## note: causal mask will not applied in cross-attention case
|
||||
x[j] = block(x[j], context=context_j)
|
||||
|
||||
if self.use_linear:
|
||||
x = self.proj_out(x)
|
||||
x = rearrange(x, 'b (h w) t c -> b c t h w', h=h, w=w).contiguous()
|
||||
if not self.use_linear:
|
||||
x = rearrange(x, 'b hw t c -> (b hw) c t').contiguous()
|
||||
x = self.proj_out(x)
|
||||
x = rearrange(x, '(b h w) c t -> b c t h w', b=b, h=h, w=w).contiguous()
|
||||
|
||||
return x + x_in
|
||||
|
||||
|
||||
class GEGLU(nn.Module):
|
||||
def __init__(self, dim_in, dim_out, device=None, dtype=None, operations=ops):
|
||||
super().__init__()
|
||||
self.proj = operations.Linear(dim_in, dim_out * 2, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x):
|
||||
x, gate = self.proj(x).chunk(2, dim=-1)
|
||||
return x * F.gelu(gate)
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0., device=None, dtype=None, operations=ops):
|
||||
super().__init__()
|
||||
inner_dim = int(dim * mult)
|
||||
dim_out = default(dim_out, dim)
|
||||
project_in = nn.Sequential(
|
||||
operations.Linear(dim, inner_dim, device=device, dtype=dtype),
|
||||
nn.GELU()
|
||||
) if not glu else GEGLU(dim, inner_dim)
|
||||
|
||||
self.net = nn.Sequential(
|
||||
project_in,
|
||||
nn.Dropout(dropout),
|
||||
operations.Linear(inner_dim, dim_out, device=device, dtype=dtype)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.net(x)
|
||||
|
||||
|
||||
class LinearAttention(nn.Module):
|
||||
def __init__(self, dim, heads=4, dim_head=32, device=None, dtype=None, operations=ops):
|
||||
super().__init__()
|
||||
self.heads = heads
|
||||
hidden_dim = dim_head * heads
|
||||
self.to_qkv = operations.Conv2d(dim, hidden_dim * 3, 1, bias = False, device=device, dtype=dtype)
|
||||
self.to_out = operations.Conv2d(hidden_dim, dim, 1, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x):
|
||||
b, c, h, w = x.shape
|
||||
qkv = self.to_qkv(x)
|
||||
q, k, v = rearrange(qkv, 'b (qkv heads c) h w -> qkv b heads c (h w)', heads = self.heads, qkv=3)
|
||||
k = k.softmax(dim=-1)
|
||||
context = torch.einsum('bhdn,bhen->bhde', k, v)
|
||||
out = torch.einsum('bhde,bhdn->bhen', context, q)
|
||||
out = rearrange(out, 'b heads c (h w) -> b (heads c) h w', heads=self.heads, h=h, w=w)
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class SpatialSelfAttention(nn.Module):
|
||||
def __init__(self, in_channels, device=None, dtype=None, operations=ops):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = operations.GroupNorm(
|
||||
num_groups=32,
|
||||
num_channels=in_channels,
|
||||
eps=1e-6,
|
||||
affine=True,
|
||||
device=device,
|
||||
dtype=dtype
|
||||
)
|
||||
self.q = operations.Conv2d(
|
||||
in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
device=device,
|
||||
dtype=dtype
|
||||
)
|
||||
self.k = operations.Conv2d(
|
||||
in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
device=device,
|
||||
dtype=dtype
|
||||
)
|
||||
self.v = operations.Conv2d(
|
||||
in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
device=device,
|
||||
dtype=dtype
|
||||
)
|
||||
self.proj_out = operations.Conv2d(
|
||||
in_channels,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
device=device,
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
h_ = x
|
||||
h_ = self.norm(h_)
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
|
||||
# compute attention
|
||||
b,c,h,w = q.shape
|
||||
q = rearrange(q, 'b c h w -> b (h w) c')
|
||||
k = rearrange(k, 'b c h w -> b c (h w)')
|
||||
w_ = torch.einsum('bij,bjk->bik', q, k)
|
||||
|
||||
w_ = w_ * (int(c)**(-0.5))
|
||||
w_ = torch.nn.functional.softmax(w_, dim=2)
|
||||
|
||||
# attend to values
|
||||
v = rearrange(v, 'b c h w -> b c (h w)')
|
||||
w_ = rearrange(w_, 'b i j -> b j i')
|
||||
h_ = torch.einsum('bij,bjk->bik', v, w_)
|
||||
h_ = rearrange(h_, 'b c (h w) -> b c h w', h=h)
|
||||
h_ = self.proj_out(h_)
|
||||
|
||||
return x+h_
|
||||
@@ -0,0 +1,389 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import kornia
|
||||
import open_clip
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
from transformers import T5Tokenizer, T5EncoderModel, CLIPTokenizer, CLIPTextModel
|
||||
from ..common import autocast
|
||||
from utils.utils import count_params
|
||||
|
||||
|
||||
class AbstractEncoder(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def encode(self, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class IdentityEncoder(AbstractEncoder):
|
||||
def encode(self, x):
|
||||
return x
|
||||
|
||||
|
||||
class ClassEmbedder(nn.Module):
|
||||
def __init__(self, embed_dim, n_classes=1000, key='class', ucg_rate=0.1):
|
||||
super().__init__()
|
||||
self.key = key
|
||||
self.embedding = nn.Embedding(n_classes, embed_dim)
|
||||
self.n_classes = n_classes
|
||||
self.ucg_rate = ucg_rate
|
||||
|
||||
def forward(self, batch, key=None, disable_dropout=False):
|
||||
if key is None:
|
||||
key = self.key
|
||||
# this is for use in crossattn
|
||||
c = batch[key][:, None]
|
||||
if self.ucg_rate > 0. and not disable_dropout:
|
||||
mask = 1. - torch.bernoulli(torch.ones_like(c) * self.ucg_rate)
|
||||
c = mask * c + (1 - mask) * torch.ones_like(c) * (self.n_classes - 1)
|
||||
c = c.long()
|
||||
c = self.embedding(c)
|
||||
return c
|
||||
|
||||
def get_unconditional_conditioning(self, bs, device="cuda"):
|
||||
uc_class = self.n_classes - 1 # 1000 classes --> 0 ... 999, one extra class for ucg (class 1000)
|
||||
uc = torch.ones((bs,), device=device) * uc_class
|
||||
uc = {self.key: uc}
|
||||
return uc
|
||||
|
||||
|
||||
def disabled_train(self, mode=True):
|
||||
"""Overwrite model.train with this function to make sure train/eval mode
|
||||
does not change anymore."""
|
||||
return self
|
||||
|
||||
|
||||
class FrozenT5Embedder(AbstractEncoder):
|
||||
"""Uses the T5 transformer encoder for text"""
|
||||
|
||||
def __init__(self, version="google/t5-v1_1-large", device="cuda", max_length=77,
|
||||
freeze=True): # others are google/t5-v1_1-xl and google/t5-v1_1-xxl
|
||||
super().__init__()
|
||||
self.tokenizer = T5Tokenizer.from_pretrained(version)
|
||||
self.transformer = T5EncoderModel.from_pretrained(version)
|
||||
self.device = device
|
||||
self.max_length = max_length # TODO: typical value?
|
||||
if freeze:
|
||||
self.freeze()
|
||||
|
||||
def freeze(self):
|
||||
self.transformer = self.transformer.eval()
|
||||
# self.train = disabled_train
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, text):
|
||||
batch_encoding = self.tokenizer(text, truncation=True, max_length=self.max_length, return_length=True,
|
||||
return_overflowing_tokens=False, padding="max_length", return_tensors="pt")
|
||||
tokens = batch_encoding["input_ids"].to(self.device)
|
||||
outputs = self.transformer(input_ids=tokens)
|
||||
|
||||
z = outputs.last_hidden_state
|
||||
return z
|
||||
|
||||
def encode(self, text):
|
||||
return self(text)
|
||||
|
||||
|
||||
class FrozenCLIPEmbedder(AbstractEncoder):
|
||||
"""Uses the CLIP transformer encoder for text (from huggingface)"""
|
||||
LAYERS = [
|
||||
"last",
|
||||
"pooled",
|
||||
"hidden"
|
||||
]
|
||||
|
||||
def __init__(self, version="openai/clip-vit-large-patch14", device="cuda", max_length=77,
|
||||
freeze=True, layer="last", layer_idx=None): # clip-vit-base-patch32
|
||||
super().__init__()
|
||||
assert layer in self.LAYERS
|
||||
self.tokenizer = CLIPTokenizer.from_pretrained(version)
|
||||
self.transformer = CLIPTextModel.from_pretrained(version)
|
||||
self.device = device
|
||||
self.max_length = max_length
|
||||
if freeze:
|
||||
self.freeze()
|
||||
self.layer = layer
|
||||
self.layer_idx = layer_idx
|
||||
if layer == "hidden":
|
||||
assert layer_idx is not None
|
||||
assert 0 <= abs(layer_idx) <= 12
|
||||
|
||||
def freeze(self):
|
||||
self.transformer = self.transformer.eval()
|
||||
# self.train = disabled_train
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, text):
|
||||
batch_encoding = self.tokenizer(text, truncation=True, max_length=self.max_length, return_length=True,
|
||||
return_overflowing_tokens=False, padding="max_length", return_tensors="pt")
|
||||
tokens = batch_encoding["input_ids"].to(self.device)
|
||||
outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden")
|
||||
if self.layer == "last":
|
||||
z = outputs.last_hidden_state
|
||||
elif self.layer == "pooled":
|
||||
z = outputs.pooler_output[:, None, :]
|
||||
else:
|
||||
z = outputs.hidden_states[self.layer_idx]
|
||||
return z
|
||||
|
||||
def encode(self, text):
|
||||
return self(text)
|
||||
|
||||
|
||||
class ClipImageEmbedder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
model,
|
||||
jit=False,
|
||||
device='cuda' if torch.cuda.is_available() else 'cpu',
|
||||
antialias=True,
|
||||
ucg_rate=0.
|
||||
):
|
||||
super().__init__()
|
||||
from clip import load as load_clip
|
||||
self.model, _ = load_clip(name=model, device=device, jit=jit)
|
||||
|
||||
self.antialias = antialias
|
||||
|
||||
self.register_buffer('mean', torch.Tensor([0.48145466, 0.4578275, 0.40821073]), persistent=False)
|
||||
self.register_buffer('std', torch.Tensor([0.26862954, 0.26130258, 0.27577711]), persistent=False)
|
||||
self.ucg_rate = ucg_rate
|
||||
|
||||
def preprocess(self, x):
|
||||
# normalize to [0,1]
|
||||
x = kornia.geometry.resize(x, (224, 224),
|
||||
interpolation='bicubic', align_corners=True,
|
||||
antialias=self.antialias)
|
||||
x = (x + 1.) / 2.
|
||||
# re-normalize according to clip
|
||||
x = kornia.enhance.normalize(x, self.mean, self.std)
|
||||
return x
|
||||
|
||||
def forward(self, x, no_dropout=False):
|
||||
# x is assumed to be in range [-1,1]
|
||||
out = self.model.encode_image(self.preprocess(x))
|
||||
out = out.to(x.dtype)
|
||||
if self.ucg_rate > 0. and not no_dropout:
|
||||
out = torch.bernoulli((1. - self.ucg_rate) * torch.ones(out.shape[0], device=out.device))[:, None] * out
|
||||
return out
|
||||
|
||||
|
||||
class FrozenOpenCLIPEmbedder(AbstractEncoder):
|
||||
"""
|
||||
Uses the OpenCLIP transformer encoder for text
|
||||
"""
|
||||
LAYERS = [
|
||||
# "pooled",
|
||||
"last",
|
||||
"penultimate"
|
||||
]
|
||||
|
||||
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77,
|
||||
freeze=True, layer="last"):
|
||||
super().__init__()
|
||||
assert layer in self.LAYERS
|
||||
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), pretrained=version)
|
||||
del model.visual
|
||||
self.model = model
|
||||
|
||||
self.device = device
|
||||
self.max_length = max_length
|
||||
if freeze:
|
||||
self.freeze()
|
||||
self.layer = layer
|
||||
if self.layer == "last":
|
||||
self.layer_idx = 0
|
||||
elif self.layer == "penultimate":
|
||||
self.layer_idx = 1
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
def freeze(self):
|
||||
self.model = self.model.eval()
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, text):
|
||||
tokens = open_clip.tokenize(text) ## all clip models use 77 as context length
|
||||
z = self.encode_with_transformer(tokens.to(self.device))
|
||||
return z
|
||||
|
||||
def encode_with_transformer(self, text):
|
||||
x = self.model.token_embedding(text) # [batch_size, n_ctx, d_model]
|
||||
x = x + self.model.positional_embedding
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.text_transformer_forward(x, attn_mask=self.model.attn_mask)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
x = self.model.ln_final(x)
|
||||
return x
|
||||
|
||||
def text_transformer_forward(self, x: torch.Tensor, attn_mask=None):
|
||||
for i, r in enumerate(self.model.transformer.resblocks):
|
||||
if i == len(self.model.transformer.resblocks) - self.layer_idx:
|
||||
break
|
||||
if self.model.transformer.grad_checkpointing and not torch.jit.is_scripting():
|
||||
x = checkpoint(r, x, attn_mask)
|
||||
else:
|
||||
x = r(x, attn_mask=attn_mask)
|
||||
return x
|
||||
|
||||
def encode(self, text):
|
||||
return self(text)
|
||||
|
||||
|
||||
class FrozenOpenCLIPImageEmbedder(AbstractEncoder):
|
||||
"""
|
||||
Uses the OpenCLIP vision transformer encoder for images
|
||||
"""
|
||||
|
||||
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77,
|
||||
freeze=True, layer="pooled", antialias=True, ucg_rate=0.):
|
||||
super().__init__()
|
||||
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'),
|
||||
pretrained=version, )
|
||||
del model.transformer
|
||||
self.model = model
|
||||
# self.mapper = torch.nn.Linear(1280, 1024)
|
||||
self.device = device
|
||||
self.max_length = max_length
|
||||
if freeze:
|
||||
self.freeze()
|
||||
self.layer = layer
|
||||
if self.layer == "penultimate":
|
||||
raise NotImplementedError()
|
||||
self.layer_idx = 1
|
||||
|
||||
self.antialias = antialias
|
||||
|
||||
self.register_buffer('mean', torch.Tensor([0.48145466, 0.4578275, 0.40821073]), persistent=False)
|
||||
self.register_buffer('std', torch.Tensor([0.26862954, 0.26130258, 0.27577711]), persistent=False)
|
||||
self.ucg_rate = ucg_rate
|
||||
|
||||
def preprocess(self, x):
|
||||
# normalize to [0,1]
|
||||
x = kornia.geometry.resize(x, (224, 224),
|
||||
interpolation='bicubic', align_corners=True,
|
||||
antialias=self.antialias)
|
||||
x = (x + 1.) / 2.
|
||||
# renormalize according to clip
|
||||
x = kornia.enhance.normalize(x, self.mean, self.std)
|
||||
return x
|
||||
|
||||
def freeze(self):
|
||||
self.model = self.model.eval()
|
||||
for param in self.model.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
@autocast
|
||||
def forward(self, image, no_dropout=False):
|
||||
z = self.encode_with_vision_transformer(image)
|
||||
if self.ucg_rate > 0. and not no_dropout:
|
||||
z = torch.bernoulli((1. - self.ucg_rate) * torch.ones(z.shape[0], device=z.device))[:, None] * z
|
||||
return z
|
||||
|
||||
def encode_with_vision_transformer(self, img):
|
||||
img = self.preprocess(img)
|
||||
x = self.model.visual(img)
|
||||
return x
|
||||
|
||||
def encode(self, text):
|
||||
return self(text)
|
||||
|
||||
class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder):
|
||||
"""
|
||||
Uses the OpenCLIP vision transformer encoder for images
|
||||
"""
|
||||
|
||||
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda",
|
||||
freeze=True, layer="pooled", antialias=True):
|
||||
super().__init__()
|
||||
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'),
|
||||
pretrained=version, )
|
||||
del model.transformer
|
||||
self.model = model
|
||||
self.device = device
|
||||
|
||||
if freeze:
|
||||
self.freeze()
|
||||
self.layer = layer
|
||||
if self.layer == "penultimate":
|
||||
raise NotImplementedError()
|
||||
self.layer_idx = 1
|
||||
|
||||
self.antialias = antialias
|
||||
|
||||
self.register_buffer('mean', torch.Tensor([0.48145466, 0.4578275, 0.40821073]), persistent=False)
|
||||
self.register_buffer('std', torch.Tensor([0.26862954, 0.26130258, 0.27577711]), persistent=False)
|
||||
|
||||
|
||||
def preprocess(self, x):
|
||||
# normalize to [0,1]
|
||||
x = kornia.geometry.resize(x, (224, 224),
|
||||
interpolation='bicubic', align_corners=True,
|
||||
antialias=self.antialias)
|
||||
x = (x + 1.) / 2.
|
||||
# renormalize according to clip
|
||||
x = kornia.enhance.normalize(x, self.mean, self.std)
|
||||
return x
|
||||
|
||||
def freeze(self):
|
||||
self.model = self.model.eval()
|
||||
for param in self.model.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, image, no_dropout=False):
|
||||
## image: b c h w
|
||||
z = self.encode_with_vision_transformer(image)
|
||||
return z
|
||||
|
||||
def encode_with_vision_transformer(self, x):
|
||||
x = self.preprocess(x)
|
||||
|
||||
# to patches - whether to use dual patchnorm - https://arxiv.org/abs/2302.01327v1
|
||||
if self.model.visual.input_patchnorm:
|
||||
# einops - rearrange(x, 'b c (h p1) (w p2) -> b (h w) (c p1 p2)')
|
||||
x = x.reshape(x.shape[0], x.shape[1], self.model.visual.grid_size[0], self.model.visual.patch_size[0], self.model.visual.grid_size[1], self.model.visual.patch_size[1])
|
||||
x = x.permute(0, 2, 4, 1, 3, 5)
|
||||
x = x.reshape(x.shape[0], self.model.visual.grid_size[0] * self.model.visual.grid_size[1], -1)
|
||||
x = self.model.visual.patchnorm_pre_ln(x)
|
||||
x = self.model.visual.conv1(x)
|
||||
else:
|
||||
x = self.model.visual.conv1(x) # shape = [*, width, grid, grid]
|
||||
x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2]
|
||||
x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width]
|
||||
|
||||
# class embeddings and positional embeddings
|
||||
x = torch.cat(
|
||||
[self.model.visual.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device),
|
||||
x], dim=1) # shape = [*, grid ** 2 + 1, width]
|
||||
x = x + self.model.visual.positional_embedding.to(x.dtype)
|
||||
|
||||
# a patch_dropout of 0. would mean it is disabled and this function would do nothing but return what was passed in
|
||||
x = self.model.visual.patch_dropout(x)
|
||||
x = self.model.visual.ln_pre(x)
|
||||
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.model.visual.transformer(x)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
|
||||
return x
|
||||
|
||||
class FrozenCLIPT5Encoder(AbstractEncoder):
|
||||
def __init__(self, clip_version="openai/clip-vit-large-patch14", t5_version="google/t5-v1_1-xl", device="cuda",
|
||||
clip_max_length=77, t5_max_length=77):
|
||||
super().__init__()
|
||||
self.clip_encoder = FrozenCLIPEmbedder(clip_version, device, max_length=clip_max_length)
|
||||
self.t5_encoder = FrozenT5Embedder(t5_version, device, max_length=t5_max_length)
|
||||
print(f"{self.clip_encoder.__class__.__name__} has {count_params(self.clip_encoder) * 1.e-6:.2f} M parameters, "
|
||||
f"{self.t5_encoder.__class__.__name__} comes with {count_params(self.t5_encoder) * 1.e-6:.2f} M params.")
|
||||
|
||||
def encode(self, text):
|
||||
return self(text)
|
||||
|
||||
def forward(self, text):
|
||||
clip_z = self.clip_encoder.encode(text)
|
||||
t5_z = self.t5_encoder.encode(text)
|
||||
return [clip_z, t5_z]
|
||||
@@ -0,0 +1,145 @@
|
||||
# modified from https://github.com/mlfoundations/open_flamingo/blob/main/open_flamingo/src/helpers.py
|
||||
# and https://github.com/lucidrains/imagen-pytorch/blob/main/imagen_pytorch/imagen_pytorch.py
|
||||
# and https://github.com/tencent-ailab/IP-Adapter/blob/main/ip_adapter/resampler.py
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class ImageProjModel(nn.Module):
|
||||
"""Projection Model"""
|
||||
def __init__(self, cross_attention_dim=1024, clip_embeddings_dim=1024, clip_extra_context_tokens=4):
|
||||
super().__init__()
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.clip_extra_context_tokens = clip_extra_context_tokens
|
||||
self.proj = nn.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim)
|
||||
self.norm = nn.LayerNorm(cross_attention_dim)
|
||||
|
||||
def forward(self, image_embeds):
|
||||
#embeds = image_embeds
|
||||
embeds = image_embeds.type(list(self.proj.parameters())[0].dtype)
|
||||
clip_extra_context_tokens = self.proj(embeds).reshape(-1, self.clip_extra_context_tokens, self.cross_attention_dim)
|
||||
clip_extra_context_tokens = self.norm(clip_extra_context_tokens)
|
||||
return clip_extra_context_tokens
|
||||
|
||||
|
||||
# FFN
|
||||
def FeedForward(dim, mult=4):
|
||||
inner_dim = int(dim * mult)
|
||||
return nn.Sequential(
|
||||
nn.LayerNorm(dim),
|
||||
nn.Linear(dim, inner_dim, bias=False),
|
||||
nn.GELU(),
|
||||
nn.Linear(inner_dim, dim, bias=False),
|
||||
)
|
||||
|
||||
|
||||
def reshape_tensor(x, heads):
|
||||
bs, length, width = x.shape
|
||||
#(bs, length, width) --> (bs, length, n_heads, dim_per_head)
|
||||
x = x.view(bs, length, heads, -1)
|
||||
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
|
||||
x = x.transpose(1, 2)
|
||||
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
|
||||
x = x.reshape(bs, heads, length, -1)
|
||||
return x
|
||||
|
||||
|
||||
class PerceiverAttention(nn.Module):
|
||||
def __init__(self, *, dim, dim_head=64, heads=8):
|
||||
super().__init__()
|
||||
self.scale = dim_head**-0.5
|
||||
self.dim_head = dim_head
|
||||
self.heads = heads
|
||||
inner_dim = dim_head * heads
|
||||
|
||||
self.norm1 = nn.LayerNorm(dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
|
||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
|
||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
||||
|
||||
|
||||
def forward(self, x, latents):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): image features
|
||||
shape (b, n1, D)
|
||||
latent (torch.Tensor): latent features
|
||||
shape (b, n2, D)
|
||||
"""
|
||||
x = self.norm1(x)
|
||||
latents = self.norm2(latents)
|
||||
|
||||
b, l, _ = latents.shape
|
||||
|
||||
q = self.to_q(latents)
|
||||
kv_input = torch.cat((x, latents), dim=-2)
|
||||
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
|
||||
|
||||
q = reshape_tensor(q, self.heads)
|
||||
k = reshape_tensor(k, self.heads)
|
||||
v = reshape_tensor(v, self.heads)
|
||||
|
||||
# attention
|
||||
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
|
||||
weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards
|
||||
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
|
||||
out = weight @ v
|
||||
|
||||
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class Resampler(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim=1024,
|
||||
depth=8,
|
||||
dim_head=64,
|
||||
heads=16,
|
||||
num_queries=8,
|
||||
embedding_dim=768,
|
||||
output_dim=1024,
|
||||
ff_mult=4,
|
||||
video_length=None, # using frame-wise version or not
|
||||
):
|
||||
super().__init__()
|
||||
## queries for a single frame / image
|
||||
self.num_queries = num_queries
|
||||
self.video_length = video_length
|
||||
|
||||
## <num_queries> queries for each frame
|
||||
if video_length is not None:
|
||||
num_queries = num_queries * video_length
|
||||
|
||||
self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5)
|
||||
self.proj_in = nn.Linear(embedding_dim, dim)
|
||||
self.proj_out = nn.Linear(dim, output_dim)
|
||||
self.norm_out = nn.LayerNorm(output_dim)
|
||||
|
||||
self.layers = nn.ModuleList([])
|
||||
for _ in range(depth):
|
||||
self.layers.append(
|
||||
nn.ModuleList(
|
||||
[
|
||||
PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
|
||||
FeedForward(dim=dim, mult=ff_mult),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
latents = self.latents.repeat(x.size(0), 1, 1) ## B (T L) C
|
||||
x = self.proj_in(x)
|
||||
|
||||
for attn, ff in self.layers:
|
||||
latents = attn(x, latents) + latents
|
||||
latents = ff(latents) + latents
|
||||
|
||||
latents = self.proj_out(latents)
|
||||
latents = self.norm_out(latents) # B L C or B (T L) C
|
||||
|
||||
return latents
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,822 @@
|
||||
from functools import partial
|
||||
from abc import abstractmethod
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
import torch.nn.functional as F
|
||||
from ...models.utils_diffusion import timestep_embedding
|
||||
from ...common import checkpoint
|
||||
from ...basics import (
|
||||
zero_module,
|
||||
conv_nd,
|
||||
linear,
|
||||
avg_pool_nd,
|
||||
normalization
|
||||
)
|
||||
from ...modules.attention import SpatialTransformer, TemporalTransformer
|
||||
import comfy.ops
|
||||
import logging
|
||||
|
||||
ops = comfy.ops.disable_weight_init
|
||||
|
||||
class TimestepBlock(nn.Module):
|
||||
"""
|
||||
Any module where forward() takes timestep embeddings as a second argument.
|
||||
"""
|
||||
@abstractmethod
|
||||
def forward(self, x, emb):
|
||||
"""
|
||||
Apply the module to `x` given `emb` timestep embeddings.
|
||||
"""
|
||||
|
||||
#This is needed because accelerate makes a copy of transformer_options which breaks "transformer_index"
|
||||
def forward_timestep_embed(ts, x, emb, context=None, batch_size=None, transformer_options={}):
|
||||
for layer in ts:
|
||||
if isinstance(layer, TimestepBlock):
|
||||
x = layer(x, emb, batch_size=batch_size)
|
||||
elif isinstance(layer, SpatialTransformer):
|
||||
x = layer(x, context)
|
||||
if "transformer_index" in transformer_options:
|
||||
transformer_options["transformer_index"] += 1
|
||||
elif isinstance(layer, TemporalTransformer):
|
||||
x = rearrange(x, '(b f) c h w -> b c f h w', b=batch_size)
|
||||
x = layer(x, context)
|
||||
if "transformer_index" in transformer_options:
|
||||
transformer_options["transformer_index"] += 1
|
||||
x = rearrange(x, 'b c f h w -> (b f) c h w')
|
||||
else:
|
||||
x = layer(x)
|
||||
return x
|
||||
|
||||
class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
|
||||
"""
|
||||
A sequential module that passes timestep embeddings to the children that
|
||||
support it as an extra input.
|
||||
"""
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
return forward_timestep_embed(self, *args, **kwargs)
|
||||
|
||||
class Downsample(nn.Module):
|
||||
"""
|
||||
A downsampling layer with an optional convolution.
|
||||
:param channels: channels in the inputs and outputs.
|
||||
:param use_conv: a bool determining if a convolution is applied.
|
||||
:param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then
|
||||
downsampling occurs in the inner-two dimensions.
|
||||
"""
|
||||
|
||||
def __init__(self, channels, use_conv, dims=2, out_channels=None, padding=1, dtype=None, device=None, operations=ops):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.dims = dims
|
||||
stride = 2 if dims != 3 else (1, 2, 2)
|
||||
if use_conv:
|
||||
self.op = operations.conv_nd(
|
||||
dims, self.channels, self.out_channels, 3, stride=stride, padding=padding
|
||||
)
|
||||
else:
|
||||
assert self.channels == self.out_channels
|
||||
self.op = avg_pool_nd(dims, kernel_size=stride, stride=stride)
|
||||
|
||||
def forward(self, x):
|
||||
assert x.shape[1] == self.channels
|
||||
return self.op(x)
|
||||
|
||||
class Upsample(nn.Module):
|
||||
"""
|
||||
An upsampling layer with an optional convolution.
|
||||
:param channels: channels in the inputs and outputs.
|
||||
:param use_conv: a bool determining if a convolution is applied.
|
||||
:param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then
|
||||
upsampling occurs in the inner-two dimensions.
|
||||
"""
|
||||
|
||||
def __init__(self, channels, use_conv, dims=2, out_channels=None, padding=1, dtype=None, device=None, operations=ops):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.dims = dims
|
||||
if use_conv:
|
||||
self.conv = operations.conv_nd(dims, self.channels, self.out_channels, 3, padding=padding, dtype=dtype, device=device)
|
||||
|
||||
def forward(self, x):
|
||||
assert x.shape[1] == self.channels
|
||||
if self.dims == 3:
|
||||
x = F.interpolate(x, (x.shape[2], x.shape[3] * 2, x.shape[4] * 2), mode='nearest')
|
||||
else:
|
||||
x = F.interpolate(x, scale_factor=2, mode='nearest')
|
||||
if self.use_conv:
|
||||
x = self.conv(x)
|
||||
return x
|
||||
|
||||
class ResBlock(TimestepBlock):
|
||||
"""
|
||||
A residual block that can optionally change the number of channels.
|
||||
:param channels: the number of input channels.
|
||||
:param emb_channels: the number of timestep embedding channels.
|
||||
:param dropout: the rate of dropout.
|
||||
:param out_channels: if specified, the number of out channels.
|
||||
:param use_conv: if True and out_channels is specified, use a spatial
|
||||
convolution instead of a smaller 1x1 convolution to change the
|
||||
channels in the skip connection.
|
||||
:param dims: determines if the signal is 1D, 2D, or 3D.
|
||||
:param up: if True, use this block for upsampling.
|
||||
:param down: if True, use this block for downsampling.
|
||||
:param use_temporal_conv: if True, use the temporal convolution.
|
||||
:param use_image_dataset: if True, the temporal parameters will not be optimized.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
channels,
|
||||
emb_channels,
|
||||
dropout,
|
||||
out_channels=None,
|
||||
use_scale_shift_norm=False,
|
||||
dims=2,
|
||||
use_checkpoint=False,
|
||||
use_conv=False,
|
||||
up=False,
|
||||
down=False,
|
||||
kernel_size=3,
|
||||
use_temporal_conv=False,
|
||||
tempspatial_aware=False,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=ops
|
||||
):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.emb_channels = emb_channels
|
||||
self.dropout = dropout
|
||||
self.out_channels = out_channels or channels
|
||||
self.use_conv = use_conv
|
||||
self.use_checkpoint = use_checkpoint
|
||||
self.use_scale_shift_norm = use_scale_shift_norm
|
||||
self.use_temporal_conv = use_temporal_conv
|
||||
|
||||
if isinstance(kernel_size, list):
|
||||
padding =[k // 2 for k in kernel_size]
|
||||
else:
|
||||
padding = kernel_size // 2
|
||||
|
||||
# operations used in normalization function
|
||||
self.in_layers = nn.Sequential(
|
||||
normalization(channels, dtype=dtype, device=device),
|
||||
nn.SiLU(),
|
||||
operations.conv_nd(dims, channels, self.out_channels, 3, padding=1, dtype=dtype, device=device),
|
||||
)
|
||||
|
||||
self.updown = up or down
|
||||
|
||||
if up:
|
||||
self.h_upd = Upsample(channels, False, dims, dtype=dtype, device=device)
|
||||
self.x_upd = Upsample(channels, False, dims, dtype=dtype, device=device)
|
||||
elif down:
|
||||
self.h_upd = Downsample(channels, False, dims, dtype=dtype, device=device)
|
||||
self.x_upd = Downsample(channels, False, dims, dtype=dtype, device=device)
|
||||
else:
|
||||
self.h_upd = self.x_upd = nn.Identity()
|
||||
|
||||
self.emb_layers = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
operations.Linear(
|
||||
emb_channels,
|
||||
2 * self.out_channels if use_scale_shift_norm else self.out_channels,
|
||||
dtype=dtype,
|
||||
device=device
|
||||
),
|
||||
)
|
||||
self.out_layers = nn.Sequential(
|
||||
normalization(self.out_channels, dtype=dtype, device=device),
|
||||
nn.SiLU(),
|
||||
nn.Dropout(p=dropout),
|
||||
zero_module(operations.Conv2d(self.out_channels, self.out_channels, 3, padding=1, dtype=dtype, device=device)),
|
||||
)
|
||||
|
||||
if self.out_channels == channels:
|
||||
self.skip_connection = nn.Identity()
|
||||
elif use_conv:
|
||||
self.skip_connection = operations.conv_nd(dims, channels, self.out_channels, 3, padding=1, dtype=dtype, device=device)
|
||||
else:
|
||||
self.skip_connection = operations.conv_nd(dims, channels, self.out_channels, 1, dtype=dtype, device=device)
|
||||
|
||||
if self.use_temporal_conv:
|
||||
self.temopral_conv = TemporalConvBlock(
|
||||
self.out_channels,
|
||||
self.out_channels,
|
||||
dropout=0.1,
|
||||
spatial_aware=tempspatial_aware,
|
||||
dtype=dtype,
|
||||
device=device
|
||||
)
|
||||
|
||||
def forward(self, x, emb, batch_size=None):
|
||||
"""
|
||||
Apply the block to a Tensor, conditioned on a timestep embedding.
|
||||
:param x: an [N x C x ...] Tensor of features.
|
||||
:param emb: an [N x emb_channels] Tensor of timestep embeddings.
|
||||
:return: an [N x C x ...] Tensor of outputs.
|
||||
"""
|
||||
input_tuple = (x, emb)
|
||||
if batch_size:
|
||||
forward_batchsize = partial(self._forward, batch_size=batch_size)
|
||||
return checkpoint(forward_batchsize, input_tuple, self.parameters(), self.use_checkpoint)
|
||||
return checkpoint(self._forward, input_tuple, self.parameters(), self.use_checkpoint)
|
||||
|
||||
def _forward(self, x, emb, batch_size=None):
|
||||
if self.updown:
|
||||
in_rest, in_conv = self.in_layers[:-1], self.in_layers[-1]
|
||||
h = in_rest(x)
|
||||
h = self.h_upd(h)
|
||||
x = self.x_upd(x)
|
||||
h = in_conv(h)
|
||||
else:
|
||||
h = self.in_layers(x)
|
||||
emb_out = self.emb_layers(emb).type(h.dtype)
|
||||
while len(emb_out.shape) < len(h.shape):
|
||||
emb_out = emb_out[..., None]
|
||||
if self.use_scale_shift_norm:
|
||||
out_norm, out_rest = self.out_layers[0], self.out_layers[1:]
|
||||
scale, shift = torch.chunk(emb_out, 2, dim=1)
|
||||
h = out_norm(h) * (1 + scale) + shift
|
||||
h = out_rest(h)
|
||||
else:
|
||||
h = h + emb_out
|
||||
h = self.out_layers(h)
|
||||
h = self.skip_connection(x) + h
|
||||
|
||||
if self.use_temporal_conv and batch_size:
|
||||
h = rearrange(h, '(b t) c h w -> b c t h w', b=batch_size)
|
||||
h = self.temopral_conv(h)
|
||||
h = rearrange(h, 'b c t h w -> (b t) c h w')
|
||||
return h
|
||||
|
||||
class TemporalConvBlock(nn.Module):
|
||||
"""
|
||||
Adapted from modelscope: https://github.com/modelscope/modelscope/blob/master/modelscope/models/multi_modal/video_synthesis/unet_sd.py
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels=None,
|
||||
dropout=0.0,
|
||||
spatial_aware=False,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=ops
|
||||
):
|
||||
super(TemporalConvBlock, self).__init__()
|
||||
if out_channels is None:
|
||||
out_channels = in_channels
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
th_kernel_shape = (3, 1, 1) if not spatial_aware else (3, 3, 1)
|
||||
th_padding_shape = (1, 0, 0) if not spatial_aware else (1, 1, 0)
|
||||
tw_kernel_shape = (3, 1, 1) if not spatial_aware else (3, 1, 3)
|
||||
tw_padding_shape = (1, 0, 0) if not spatial_aware else (1, 0, 1)
|
||||
|
||||
# conv layers
|
||||
self.conv1 = nn.Sequential(
|
||||
operations.GroupNorm(32, in_channels, device=device, dtype=dtype), nn.SiLU(),
|
||||
operations.Conv3d(in_channels, out_channels, th_kernel_shape, padding=th_padding_shape, device=device, dtype=dtype))
|
||||
self.conv2 = nn.Sequential(
|
||||
operations.GroupNorm(32, out_channels, device=device, dtype=dtype), nn.SiLU(), nn.Dropout(dropout),
|
||||
operations.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape, device=device, dtype=dtype))
|
||||
self.conv3 = nn.Sequential(
|
||||
operations.GroupNorm(32, out_channels, device=device, dtype=dtype), nn.SiLU(), nn.Dropout(dropout),
|
||||
operations.Conv3d(out_channels, in_channels, th_kernel_shape, padding=th_padding_shape, device=device, dtype=dtype))
|
||||
self.conv4 = nn.Sequential(
|
||||
operations.GroupNorm(32, out_channels, device=device, dtype=dtype), nn.SiLU(), nn.Dropout(dropout),
|
||||
operations.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape, device=device, dtype=dtype))
|
||||
|
||||
# zero out the last layer params,so the conv block is identity
|
||||
nn.init.zeros_(self.conv4[-1].weight)
|
||||
nn.init.zeros_(self.conv4[-1].bias)
|
||||
|
||||
def forward(self, x):
|
||||
identity = x
|
||||
x = self.conv1(x)
|
||||
x = self.conv2(x)
|
||||
x = self.conv3(x)
|
||||
x = self.conv4(x)
|
||||
|
||||
return identity + x
|
||||
|
||||
def context_processor(context, t, img_emb=None, temporal_size=16, concat_only=False, disable_concat=False):
|
||||
if disable_concat:
|
||||
return context
|
||||
|
||||
## repeat t times for context [(b t) 77 768] & time embedding
|
||||
## check if we use per-frame image conditioning
|
||||
|
||||
if img_emb is not None:
|
||||
context = torch.cat([context, img_emb.to(context.device, context.dtype)], dim=1)
|
||||
|
||||
if concat_only:
|
||||
return context
|
||||
|
||||
b, l_context, _ = context.shape
|
||||
if l_context == 77 + t * temporal_size:
|
||||
context_text, context_img = context[:,:77,:], context[:,77:,:]
|
||||
context_text = context_text.repeat_interleave(repeats=t, dim=0)
|
||||
context_img = rearrange(context_img, 'b (t l) c -> (b t) l c', t=t)
|
||||
context = torch.cat([context_text, context_img], dim=1)
|
||||
else:
|
||||
context = context.repeat_interleave(repeats=t, dim=0)
|
||||
|
||||
return context
|
||||
|
||||
def apply_control(h, control, name, cond_idx=None):
|
||||
if control is not None and name in control and len(control[name]) > 0:
|
||||
frames = h.shape[0]
|
||||
ctrl = control[name].pop()
|
||||
if ctrl is not None:
|
||||
try:
|
||||
if cond_idx is not None and ctrl.shape[0] > frames:
|
||||
ctrl_frames_list = list(range(ctrl.shape[0]))
|
||||
ctrl_frames = len(ctrl_frames_list)
|
||||
|
||||
idxs = (
|
||||
ctrl_frames_list[ctrl_frames // 2:] if cond_idx == 0 else \
|
||||
ctrl_frames_list[:ctrl_frames // 2]
|
||||
)
|
||||
|
||||
ctrl = ctrl[idxs]
|
||||
|
||||
h += ctrl
|
||||
except Exception as e:
|
||||
if h.shape != ctrl.shape:
|
||||
logging.warning(
|
||||
"warning control could not be applied {} {}".format(h.shape, ctrl.shape)
|
||||
)
|
||||
logging.warning(e)
|
||||
return h
|
||||
|
||||
class UNetModel(nn.Module):
|
||||
"""
|
||||
The full UNet model with attention and timestep embedding.
|
||||
:param in_channels: in_channels in the input Tensor.
|
||||
:param model_channels: base channel count for the model.
|
||||
:param out_channels: channels in the output Tensor.
|
||||
:param num_res_blocks: number of residual blocks per downsample.
|
||||
:param attention_resolutions: a collection of downsample rates at which
|
||||
attention will take place. May be a set, list, or tuple.
|
||||
For example, if this contains 4, then at 4x downsampling, attention
|
||||
will be used.
|
||||
:param dropout: the dropout probability.
|
||||
:param channel_mult: channel multiplier for each level of the UNet.
|
||||
:param conv_resample: if True, use learned convolutions for upsampling and
|
||||
downsampling.
|
||||
:param dims: determines if the signal is 1D, 2D, or 3D.
|
||||
:param num_classes: if specified (as an int), then this model will be
|
||||
class-conditional with `num_classes` classes.
|
||||
:param use_checkpoint: use gradient checkpointing to reduce memory usage.
|
||||
:param num_heads: the number of attention heads in each attention layer.
|
||||
:param num_heads_channels: if specified, ignore num_heads and instead use
|
||||
a fixed channel width per attention head.
|
||||
:param num_heads_upsample: works with num_heads to set a different number
|
||||
of heads for upsampling. Deprecated.
|
||||
:param use_scale_shift_norm: use a FiLM-like conditioning mechanism.
|
||||
:param resblock_updown: use residual blocks for up/downsampling.
|
||||
:param use_new_attention_order: use a different attention pattern for potentially
|
||||
increased efficiency.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
model_channels,
|
||||
out_channels,
|
||||
num_res_blocks,
|
||||
attention_resolutions,
|
||||
dropout=0.0,
|
||||
channel_mult=(1, 2, 4, 8),
|
||||
conv_resample=True,
|
||||
dims=2,
|
||||
context_dim=None,
|
||||
use_scale_shift_norm=False,
|
||||
resblock_updown=False,
|
||||
num_heads=-1,
|
||||
num_head_channels=-1,
|
||||
transformer_depth=1,
|
||||
use_linear=False,
|
||||
use_checkpoint=False,
|
||||
temporal_conv=False,
|
||||
tempspatial_aware=False,
|
||||
temporal_attention=True,
|
||||
use_relative_position=True,
|
||||
use_causal_attention=False,
|
||||
temporal_length=None,
|
||||
use_fp16=False,
|
||||
addition_attention=False,
|
||||
temporal_selfatt_only=True,
|
||||
image_cross_attention=False,
|
||||
image_cross_attention_scale_learnable=False,
|
||||
default_fs=4,
|
||||
fs_condition=False,
|
||||
device=None,
|
||||
dtype=torch.float16,
|
||||
operations=ops
|
||||
):
|
||||
super(UNetModel, self).__init__()
|
||||
if num_heads == -1:
|
||||
assert num_head_channels != -1, 'Either num_heads or num_head_channels has to be set'
|
||||
if num_head_channels == -1:
|
||||
assert num_heads != -1, 'Either num_heads or num_head_channels has to be set'
|
||||
|
||||
self.in_channels = in_channels
|
||||
self.model_channels = model_channels
|
||||
self.out_channels = out_channels
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.attention_resolutions = attention_resolutions
|
||||
self.dropout = dropout
|
||||
self.channel_mult = channel_mult
|
||||
self.conv_resample = conv_resample
|
||||
self.temporal_attention = temporal_attention
|
||||
time_embed_dim = model_channels * 4
|
||||
self.use_checkpoint = use_checkpoint
|
||||
temporal_self_att_only = True
|
||||
self.addition_attention = addition_attention
|
||||
self.temporal_length = temporal_length
|
||||
self.image_cross_attention = image_cross_attention
|
||||
self.image_cross_attention_scale_learnable = image_cross_attention_scale_learnable
|
||||
self.default_fs = default_fs
|
||||
self.fs_condition = fs_condition
|
||||
self.device = device
|
||||
#self.dtype = dtype
|
||||
self.dtype = torch.float32
|
||||
|
||||
## Time embedding blocks
|
||||
self.time_embed = nn.Sequential(
|
||||
linear(model_channels, time_embed_dim, device=device, dtype=self.dtype),
|
||||
nn.SiLU(),
|
||||
linear(time_embed_dim, time_embed_dim, device=device, dtype=self.dtype),
|
||||
)
|
||||
if fs_condition:
|
||||
self.fps_embedding = nn.Sequential(
|
||||
linear(model_channels, time_embed_dim, device=device, dtype=self.dtype),
|
||||
nn.SiLU(),
|
||||
linear(time_embed_dim, time_embed_dim, device=device, dtype=self.dtype),
|
||||
)
|
||||
nn.init.zeros_(self.fps_embedding[-1].weight)
|
||||
nn.init.zeros_(self.fps_embedding[-1].bias)
|
||||
## Input Block
|
||||
self.input_blocks = nn.ModuleList(
|
||||
[
|
||||
TimestepEmbedSequential(
|
||||
operations.conv_nd(
|
||||
dims,
|
||||
in_channels,
|
||||
model_channels,
|
||||
3,
|
||||
padding=1,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
))
|
||||
]
|
||||
)
|
||||
if self.addition_attention:
|
||||
self.init_attn=TimestepEmbedSequential(
|
||||
TemporalTransformer(
|
||||
model_channels,
|
||||
n_heads=8,
|
||||
d_head=num_head_channels,
|
||||
depth=transformer_depth,
|
||||
context_dim=context_dim,
|
||||
use_checkpoint=use_checkpoint, only_self_att=temporal_selfatt_only,
|
||||
causal_attention=False, relative_position=use_relative_position,
|
||||
temporal_length=temporal_length,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
))
|
||||
|
||||
input_block_chans = [model_channels]
|
||||
ch = model_channels
|
||||
ds = 1
|
||||
for level, mult in enumerate(channel_mult):
|
||||
for _ in range(num_res_blocks):
|
||||
layers = [
|
||||
ResBlock(ch, time_embed_dim, dropout,
|
||||
out_channels=mult * model_channels, dims=dims, use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware,
|
||||
use_temporal_conv=temporal_conv,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
]
|
||||
ch = mult * model_channels
|
||||
if ds in attention_resolutions:
|
||||
if num_head_channels == -1:
|
||||
dim_head = ch // num_heads
|
||||
else:
|
||||
num_heads = ch // num_head_channels
|
||||
dim_head = num_head_channels
|
||||
layers.append(
|
||||
SpatialTransformer(ch, num_heads, dim_head,
|
||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
||||
use_checkpoint=use_checkpoint, disable_self_attn=False,
|
||||
video_length=temporal_length, image_cross_attention=self.image_cross_attention,
|
||||
image_cross_attention_scale_learnable=self.image_cross_attention_scale_learnable,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
)
|
||||
if self.temporal_attention:
|
||||
layers.append(
|
||||
TemporalTransformer(ch, num_heads, dim_head,
|
||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
||||
use_checkpoint=use_checkpoint, only_self_att=temporal_self_att_only,
|
||||
causal_attention=use_causal_attention, relative_position=use_relative_position,
|
||||
temporal_length=temporal_length,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
)
|
||||
self.input_blocks.append(TimestepEmbedSequential(*layers))
|
||||
input_block_chans.append(ch)
|
||||
if level != len(channel_mult) - 1:
|
||||
out_ch = ch
|
||||
self.input_blocks.append(
|
||||
TimestepEmbedSequential(
|
||||
ResBlock(ch, time_embed_dim, dropout,
|
||||
out_channels=out_ch, dims=dims, use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
down=True,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
if resblock_updown
|
||||
else Downsample(
|
||||
ch,
|
||||
conv_resample,
|
||||
dims=dims,
|
||||
out_channels=out_ch,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
)
|
||||
)
|
||||
ch = out_ch
|
||||
input_block_chans.append(ch)
|
||||
ds *= 2
|
||||
|
||||
if num_head_channels == -1:
|
||||
dim_head = ch // num_heads
|
||||
else:
|
||||
num_heads = ch // num_head_channels
|
||||
dim_head = num_head_channels
|
||||
layers = [
|
||||
ResBlock(ch, time_embed_dim, dropout,
|
||||
dims=dims, use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware,
|
||||
use_temporal_conv=temporal_conv,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
),
|
||||
SpatialTransformer(ch, num_heads, dim_head,
|
||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
||||
use_checkpoint=use_checkpoint, disable_self_attn=False, video_length=temporal_length,
|
||||
image_cross_attention=self.image_cross_attention,image_cross_attention_scale_learnable=self.image_cross_attention_scale_learnable,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
]
|
||||
if self.temporal_attention:
|
||||
layers.append(
|
||||
TemporalTransformer(ch, num_heads, dim_head,
|
||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
||||
use_checkpoint=use_checkpoint, only_self_att=temporal_self_att_only,
|
||||
causal_attention=use_causal_attention, relative_position=use_relative_position,
|
||||
temporal_length=temporal_length,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
)
|
||||
layers.append(
|
||||
ResBlock(ch, time_embed_dim, dropout,
|
||||
dims=dims, use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware,
|
||||
use_temporal_conv=temporal_conv,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
)
|
||||
|
||||
## Middle Block
|
||||
self.middle_block = TimestepEmbedSequential(*layers)
|
||||
|
||||
## Output Block
|
||||
self.output_blocks = nn.ModuleList([])
|
||||
for level, mult in list(enumerate(channel_mult))[::-1]:
|
||||
for i in range(num_res_blocks + 1):
|
||||
ich = input_block_chans.pop()
|
||||
layers = [
|
||||
ResBlock(ch + ich, time_embed_dim, dropout,
|
||||
out_channels=mult * model_channels, dims=dims, use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware,
|
||||
use_temporal_conv=temporal_conv,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
]
|
||||
ch = model_channels * mult
|
||||
if ds in attention_resolutions:
|
||||
if num_head_channels == -1:
|
||||
dim_head = ch // num_heads
|
||||
else:
|
||||
num_heads = ch // num_head_channels
|
||||
dim_head = num_head_channels
|
||||
layers.append(
|
||||
SpatialTransformer(ch, num_heads, dim_head,
|
||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
||||
use_checkpoint=use_checkpoint, disable_self_attn=False, video_length=temporal_length,
|
||||
image_cross_attention=self.image_cross_attention,image_cross_attention_scale_learnable=self.image_cross_attention_scale_learnable,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
)
|
||||
if self.temporal_attention:
|
||||
layers.append(
|
||||
TemporalTransformer(ch, num_heads, dim_head,
|
||||
depth=transformer_depth, context_dim=context_dim, use_linear=use_linear,
|
||||
use_checkpoint=use_checkpoint, only_self_att=temporal_self_att_only,
|
||||
causal_attention=use_causal_attention, relative_position=use_relative_position,
|
||||
temporal_length=temporal_length,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
)
|
||||
if level and i == num_res_blocks:
|
||||
out_ch = ch
|
||||
layers.append(
|
||||
ResBlock(ch, time_embed_dim, dropout,
|
||||
out_channels=out_ch, dims=dims, use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
up=True,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
if resblock_updown
|
||||
else Upsample(ch, conv_resample, dims=dims, out_channels=out_ch)
|
||||
)
|
||||
ds //= 2
|
||||
self.output_blocks.append(TimestepEmbedSequential(*layers))
|
||||
|
||||
self.out = nn.Sequential(
|
||||
normalization(ch, device=device, dtype=self.dtype),
|
||||
nn.SiLU(),
|
||||
zero_module(
|
||||
operations.conv_nd(
|
||||
dims,
|
||||
model_channels,
|
||||
out_channels,
|
||||
3,
|
||||
padding=1,
|
||||
device=device,
|
||||
dtype=self.dtype
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
# TODO Add Transformer options to leverage the usage of patches.
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
timesteps,
|
||||
context=None,
|
||||
context_in=None,
|
||||
cc_concat=None,
|
||||
num_video_frames=16,
|
||||
features_adapter=None,
|
||||
fs=None,
|
||||
img_emb=None,
|
||||
control=None,
|
||||
transformer_options={},
|
||||
cond_idx=None,
|
||||
**kwargs
|
||||
):
|
||||
|
||||
if any([fs is None, img_emb is None, cc_concat is None]):
|
||||
raise ValueError("One or more of the required inputs for UNet Forward is None.")
|
||||
|
||||
cond_idx = transformer_options.get("cond_idx", None)
|
||||
transformer_options['original_shape'] = list(x.shape)
|
||||
transformer_options['transformer_index'] = 0
|
||||
transformer_patches = transformer_options.get("patches", {})
|
||||
|
||||
# In ComfyUI, the frames are always with the batch, so we deconstruct it here.
|
||||
# This is mandatory as this is a video based model.
|
||||
# We usually denote "f" as frames, but will use "t" (time) to be consistent with DynamiCrafter.
|
||||
b,_,t,_,_ = x.shape
|
||||
|
||||
context = context_in
|
||||
cc_concat = cc_concat.to(x.device, x.dtype)
|
||||
x = torch.cat([x, cc_concat], dim=1)
|
||||
|
||||
fs = fs.to(x.device, x.dtype)
|
||||
|
||||
timestep = timesteps
|
||||
context = context_processor(context, num_video_frames, img_emb=img_emb)
|
||||
|
||||
t_emb = timestep_embedding(timestep, self.model_channels, repeat_only=False, dtype=self.dtype)
|
||||
emb = self.time_embed(t_emb)
|
||||
emb = emb.repeat_interleave(repeats=t, dim=0)
|
||||
|
||||
## always in shape (b t) c h w, except for temporal layer
|
||||
x = rearrange(x, 'b c t h w -> (b t) c h w')
|
||||
|
||||
## combine emb
|
||||
if self.fs_condition:
|
||||
if fs is None:
|
||||
fs = torch.tensor(
|
||||
[self.default_fs] * b, dtype=torch.long, device=x.device)
|
||||
fs_emb = timestep_embedding(fs, self.model_channels, repeat_only=False, dtype=self.dtype).type(x.dtype)
|
||||
|
||||
fs_embed = self.fps_embedding(fs_emb)
|
||||
fs_embed = fs_embed.repeat_interleave(repeats=t, dim=0)
|
||||
|
||||
emb = emb + fs_embed
|
||||
|
||||
h = x.type(self.dtype)
|
||||
adapter_idx = 0
|
||||
hs = []
|
||||
|
||||
for id, module in enumerate(self.input_blocks):
|
||||
transformer_options["block"] = ("input", id)
|
||||
#h = module(h, emb, context=context, batch_size=b)
|
||||
h = forward_timestep_embed(
|
||||
module,
|
||||
h,
|
||||
emb,
|
||||
context=context,
|
||||
batch_size=b,
|
||||
transformer_options=transformer_options
|
||||
)
|
||||
h = apply_control(h, control, 'input', cond_idx)
|
||||
|
||||
if "input_block_patch" in transformer_patches:
|
||||
patch = transformer_patches["input_block_patch"]
|
||||
for p in patch:
|
||||
h = p(h, transformer_options)
|
||||
|
||||
if id ==0 and self.addition_attention:
|
||||
h = forward_timestep_embed(
|
||||
self.init_attn,
|
||||
h,
|
||||
emb,
|
||||
context=context,
|
||||
batch_size=b,
|
||||
transformer_options=transformer_options
|
||||
)
|
||||
## plug-in adapter features
|
||||
if ((id+1)%3 == 0) and features_adapter is not None:
|
||||
h = h + features_adapter[adapter_idx]
|
||||
adapter_idx += 1
|
||||
hs.append(h)
|
||||
if "input_block_patch_after_skip" in transformer_patches:
|
||||
patch = transformer_patches["input_block_patch_after_skip"]
|
||||
for p in patch:
|
||||
h = p(h, transformer_options)
|
||||
if features_adapter is not None:
|
||||
assert len(features_adapter)==adapter_idx, 'Wrong features_adapter'
|
||||
transformer_options["block"] = ("middle", 0)
|
||||
h = forward_timestep_embed(
|
||||
self.middle_block,
|
||||
h,
|
||||
emb,
|
||||
context=context,
|
||||
batch_size=b,
|
||||
transformer_options=transformer_options
|
||||
)
|
||||
h = apply_control(h, control, 'middle', cond_idx)
|
||||
for id, module in enumerate(self.output_blocks):
|
||||
transformer_options["block"] = ("output", id)
|
||||
hsp = hs.pop()
|
||||
hsp = apply_control(hsp, control, 'output', cond_idx)
|
||||
|
||||
if "output_block_patch" in transformer_patches:
|
||||
patch = transformer_patches["output_block_patch"]
|
||||
for p in patch:
|
||||
h, hsp = p(h, hsp, transformer_options)
|
||||
|
||||
h = torch.cat([h, hsp], dim=1)
|
||||
del hsp
|
||||
h = forward_timestep_embed(
|
||||
module,
|
||||
h,
|
||||
emb,
|
||||
context=context,
|
||||
batch_size=b,
|
||||
transformer_options=transformer_options
|
||||
)
|
||||
h = h.type(x.dtype)
|
||||
h = self.out(h)
|
||||
|
||||
# We output with the tensor unfolded framewise, then reshape them to batched using ComfyUI nodes.
|
||||
h = rearrange(h, '(b t) c h w -> b c t h w', t=num_video_frames)
|
||||
|
||||
return h
|
||||
@@ -0,0 +1,639 @@
|
||||
"""shout-out to https://github.com/lucidrains/x-transformers/tree/main/x_transformers"""
|
||||
from functools import partial
|
||||
from inspect import isfunction
|
||||
from collections import namedtuple
|
||||
from einops import rearrange, repeat
|
||||
import torch
|
||||
from torch import nn, einsum
|
||||
import torch.nn.functional as F
|
||||
|
||||
# constants
|
||||
DEFAULT_DIM_HEAD = 64
|
||||
|
||||
Intermediates = namedtuple('Intermediates', [
|
||||
'pre_softmax_attn',
|
||||
'post_softmax_attn'
|
||||
])
|
||||
|
||||
LayerIntermediates = namedtuple('Intermediates', [
|
||||
'hiddens',
|
||||
'attn_intermediates'
|
||||
])
|
||||
|
||||
|
||||
class AbsolutePositionalEmbedding(nn.Module):
|
||||
def __init__(self, dim, max_seq_len):
|
||||
super().__init__()
|
||||
self.emb = nn.Embedding(max_seq_len, dim)
|
||||
self.init_()
|
||||
|
||||
def init_(self):
|
||||
nn.init.normal_(self.emb.weight, std=0.02)
|
||||
|
||||
def forward(self, x):
|
||||
n = torch.arange(x.shape[1], device=x.device)
|
||||
return self.emb(n)[None, :, :]
|
||||
|
||||
|
||||
class FixedPositionalEmbedding(nn.Module):
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
inv_freq = 1. / (10000 ** (torch.arange(0, dim, 2).float() / dim))
|
||||
self.register_buffer('inv_freq', inv_freq)
|
||||
|
||||
def forward(self, x, seq_dim=1, offset=0):
|
||||
t = torch.arange(x.shape[seq_dim], device=x.device).type_as(self.inv_freq) + offset
|
||||
sinusoid_inp = torch.einsum('i , j -> i j', t, self.inv_freq)
|
||||
emb = torch.cat((sinusoid_inp.sin(), sinusoid_inp.cos()), dim=-1)
|
||||
return emb[None, :, :]
|
||||
|
||||
|
||||
# helpers
|
||||
|
||||
def exists(val):
|
||||
return val is not None
|
||||
|
||||
|
||||
def default(val, d):
|
||||
if exists(val):
|
||||
return val
|
||||
return d() if isfunction(d) else d
|
||||
|
||||
|
||||
def always(val):
|
||||
def inner(*args, **kwargs):
|
||||
return val
|
||||
return inner
|
||||
|
||||
|
||||
def not_equals(val):
|
||||
def inner(x):
|
||||
return x != val
|
||||
return inner
|
||||
|
||||
|
||||
def equals(val):
|
||||
def inner(x):
|
||||
return x == val
|
||||
return inner
|
||||
|
||||
|
||||
def max_neg_value(tensor):
|
||||
return -torch.finfo(tensor.dtype).max
|
||||
|
||||
|
||||
# keyword argument helpers
|
||||
|
||||
def pick_and_pop(keys, d):
|
||||
values = list(map(lambda key: d.pop(key), keys))
|
||||
return dict(zip(keys, values))
|
||||
|
||||
|
||||
def group_dict_by_key(cond, d):
|
||||
return_val = [dict(), dict()]
|
||||
for key in d.keys():
|
||||
match = bool(cond(key))
|
||||
ind = int(not match)
|
||||
return_val[ind][key] = d[key]
|
||||
return (*return_val,)
|
||||
|
||||
|
||||
def string_begins_with(prefix, str):
|
||||
return str.startswith(prefix)
|
||||
|
||||
|
||||
def group_by_key_prefix(prefix, d):
|
||||
return group_dict_by_key(partial(string_begins_with, prefix), d)
|
||||
|
||||
|
||||
def groupby_prefix_and_trim(prefix, d):
|
||||
kwargs_with_prefix, kwargs = group_dict_by_key(partial(string_begins_with, prefix), d)
|
||||
kwargs_without_prefix = dict(map(lambda x: (x[0][len(prefix):], x[1]), tuple(kwargs_with_prefix.items())))
|
||||
return kwargs_without_prefix, kwargs
|
||||
|
||||
|
||||
# classes
|
||||
class Scale(nn.Module):
|
||||
def __init__(self, value, fn):
|
||||
super().__init__()
|
||||
self.value = value
|
||||
self.fn = fn
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
x, *rest = self.fn(x, **kwargs)
|
||||
return (x * self.value, *rest)
|
||||
|
||||
|
||||
class Rezero(nn.Module):
|
||||
def __init__(self, fn):
|
||||
super().__init__()
|
||||
self.fn = fn
|
||||
self.g = nn.Parameter(torch.zeros(1))
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
x, *rest = self.fn(x, **kwargs)
|
||||
return (x * self.g, *rest)
|
||||
|
||||
|
||||
class ScaleNorm(nn.Module):
|
||||
def __init__(self, dim, eps=1e-5):
|
||||
super().__init__()
|
||||
self.scale = dim ** -0.5
|
||||
self.eps = eps
|
||||
self.g = nn.Parameter(torch.ones(1))
|
||||
|
||||
def forward(self, x):
|
||||
norm = torch.norm(x, dim=-1, keepdim=True) * self.scale
|
||||
return x / norm.clamp(min=self.eps) * self.g
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim, eps=1e-8):
|
||||
super().__init__()
|
||||
self.scale = dim ** -0.5
|
||||
self.eps = eps
|
||||
self.g = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
norm = torch.norm(x, dim=-1, keepdim=True) * self.scale
|
||||
return x / norm.clamp(min=self.eps) * self.g
|
||||
|
||||
|
||||
class Residual(nn.Module):
|
||||
def forward(self, x, residual):
|
||||
return x + residual
|
||||
|
||||
|
||||
class GRUGating(nn.Module):
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.gru = nn.GRUCell(dim, dim)
|
||||
|
||||
def forward(self, x, residual):
|
||||
gated_output = self.gru(
|
||||
rearrange(x, 'b n d -> (b n) d'),
|
||||
rearrange(residual, 'b n d -> (b n) d')
|
||||
)
|
||||
|
||||
return gated_output.reshape_as(x)
|
||||
|
||||
|
||||
# feedforward
|
||||
|
||||
class GEGLU(nn.Module):
|
||||
def __init__(self, dim_in, dim_out):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(dim_in, dim_out * 2)
|
||||
|
||||
def forward(self, x):
|
||||
x, gate = self.proj(x).chunk(2, dim=-1)
|
||||
return x * F.gelu(gate)
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0.):
|
||||
super().__init__()
|
||||
inner_dim = int(dim * mult)
|
||||
dim_out = default(dim_out, dim)
|
||||
project_in = nn.Sequential(
|
||||
nn.Linear(dim, inner_dim),
|
||||
nn.GELU()
|
||||
) if not glu else GEGLU(dim, inner_dim)
|
||||
|
||||
self.net = nn.Sequential(
|
||||
project_in,
|
||||
nn.Dropout(dropout),
|
||||
nn.Linear(inner_dim, dim_out)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.net(x)
|
||||
|
||||
|
||||
# attention.
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
dim_head=DEFAULT_DIM_HEAD,
|
||||
heads=8,
|
||||
causal=False,
|
||||
mask=None,
|
||||
talking_heads=False,
|
||||
sparse_topk=None,
|
||||
use_entmax15=False,
|
||||
num_mem_kv=0,
|
||||
dropout=0.,
|
||||
on_attn=False
|
||||
):
|
||||
super().__init__()
|
||||
if use_entmax15:
|
||||
raise NotImplementedError("Check out entmax activation instead of softmax activation!")
|
||||
self.scale = dim_head ** -0.5
|
||||
self.heads = heads
|
||||
self.causal = causal
|
||||
self.mask = mask
|
||||
|
||||
inner_dim = dim_head * heads
|
||||
|
||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.to_k = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.to_v = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
# talking heads
|
||||
self.talking_heads = talking_heads
|
||||
if talking_heads:
|
||||
self.pre_softmax_proj = nn.Parameter(torch.randn(heads, heads))
|
||||
self.post_softmax_proj = nn.Parameter(torch.randn(heads, heads))
|
||||
|
||||
# explicit topk sparse attention
|
||||
self.sparse_topk = sparse_topk
|
||||
|
||||
# entmax
|
||||
#self.attn_fn = entmax15 if use_entmax15 else F.softmax
|
||||
self.attn_fn = F.softmax
|
||||
|
||||
# add memory key / values
|
||||
self.num_mem_kv = num_mem_kv
|
||||
if num_mem_kv > 0:
|
||||
self.mem_k = nn.Parameter(torch.randn(heads, num_mem_kv, dim_head))
|
||||
self.mem_v = nn.Parameter(torch.randn(heads, num_mem_kv, dim_head))
|
||||
|
||||
# attention on attention
|
||||
self.attn_on_attn = on_attn
|
||||
self.to_out = nn.Sequential(nn.Linear(inner_dim, dim * 2), nn.GLU()) if on_attn else nn.Linear(inner_dim, dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
context=None,
|
||||
mask=None,
|
||||
context_mask=None,
|
||||
rel_pos=None,
|
||||
sinusoidal_emb=None,
|
||||
prev_attn=None,
|
||||
mem=None
|
||||
):
|
||||
b, n, _, h, talking_heads, device = *x.shape, self.heads, self.talking_heads, x.device
|
||||
kv_input = default(context, x)
|
||||
|
||||
q_input = x
|
||||
k_input = kv_input
|
||||
v_input = kv_input
|
||||
|
||||
if exists(mem):
|
||||
k_input = torch.cat((mem, k_input), dim=-2)
|
||||
v_input = torch.cat((mem, v_input), dim=-2)
|
||||
|
||||
if exists(sinusoidal_emb):
|
||||
# in shortformer, the query would start at a position offset depending on the past cached memory
|
||||
offset = k_input.shape[-2] - q_input.shape[-2]
|
||||
q_input = q_input + sinusoidal_emb(q_input, offset=offset)
|
||||
k_input = k_input + sinusoidal_emb(k_input)
|
||||
|
||||
q = self.to_q(q_input)
|
||||
k = self.to_k(k_input)
|
||||
v = self.to_v(v_input)
|
||||
|
||||
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h=h), (q, k, v))
|
||||
|
||||
input_mask = None
|
||||
if any(map(exists, (mask, context_mask))):
|
||||
q_mask = default(mask, lambda: torch.ones((b, n), device=device).bool())
|
||||
k_mask = q_mask if not exists(context) else context_mask
|
||||
k_mask = default(k_mask, lambda: torch.ones((b, k.shape[-2]), device=device).bool())
|
||||
q_mask = rearrange(q_mask, 'b i -> b () i ()')
|
||||
k_mask = rearrange(k_mask, 'b j -> b () () j')
|
||||
input_mask = q_mask * k_mask
|
||||
|
||||
if self.num_mem_kv > 0:
|
||||
mem_k, mem_v = map(lambda t: repeat(t, 'h n d -> b h n d', b=b), (self.mem_k, self.mem_v))
|
||||
k = torch.cat((mem_k, k), dim=-2)
|
||||
v = torch.cat((mem_v, v), dim=-2)
|
||||
if exists(input_mask):
|
||||
input_mask = F.pad(input_mask, (self.num_mem_kv, 0), value=True)
|
||||
|
||||
dots = einsum('b h i d, b h j d -> b h i j', q, k) * self.scale
|
||||
mask_value = max_neg_value(dots)
|
||||
|
||||
if exists(prev_attn):
|
||||
dots = dots + prev_attn
|
||||
|
||||
pre_softmax_attn = dots
|
||||
|
||||
if talking_heads:
|
||||
dots = einsum('b h i j, h k -> b k i j', dots, self.pre_softmax_proj).contiguous()
|
||||
|
||||
if exists(rel_pos):
|
||||
dots = rel_pos(dots)
|
||||
|
||||
if exists(input_mask):
|
||||
dots.masked_fill_(~input_mask, mask_value)
|
||||
del input_mask
|
||||
|
||||
if self.causal:
|
||||
i, j = dots.shape[-2:]
|
||||
r = torch.arange(i, device=device)
|
||||
mask = rearrange(r, 'i -> () () i ()') < rearrange(r, 'j -> () () () j')
|
||||
mask = F.pad(mask, (j - i, 0), value=False)
|
||||
dots.masked_fill_(mask, mask_value)
|
||||
del mask
|
||||
|
||||
if exists(self.sparse_topk) and self.sparse_topk < dots.shape[-1]:
|
||||
top, _ = dots.topk(self.sparse_topk, dim=-1)
|
||||
vk = top[..., -1].unsqueeze(-1).expand_as(dots)
|
||||
mask = dots < vk
|
||||
dots.masked_fill_(mask, mask_value)
|
||||
del mask
|
||||
|
||||
attn = self.attn_fn(dots, dim=-1)
|
||||
post_softmax_attn = attn
|
||||
|
||||
attn = self.dropout(attn)
|
||||
|
||||
if talking_heads:
|
||||
attn = einsum('b h i j, h k -> b k i j', attn, self.post_softmax_proj).contiguous()
|
||||
|
||||
out = einsum('b h i j, b h j d -> b h i d', attn, v)
|
||||
out = rearrange(out, 'b h n d -> b n (h d)')
|
||||
|
||||
intermediates = Intermediates(
|
||||
pre_softmax_attn=pre_softmax_attn,
|
||||
post_softmax_attn=post_softmax_attn
|
||||
)
|
||||
|
||||
return self.to_out(out), intermediates
|
||||
|
||||
|
||||
class AttentionLayers(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
depth,
|
||||
heads=8,
|
||||
causal=False,
|
||||
cross_attend=False,
|
||||
only_cross=False,
|
||||
use_scalenorm=False,
|
||||
use_rmsnorm=False,
|
||||
use_rezero=False,
|
||||
rel_pos_num_buckets=32,
|
||||
rel_pos_max_distance=128,
|
||||
position_infused_attn=False,
|
||||
custom_layers=None,
|
||||
sandwich_coef=None,
|
||||
par_ratio=None,
|
||||
residual_attn=False,
|
||||
cross_residual_attn=False,
|
||||
macaron=False,
|
||||
pre_norm=True,
|
||||
gate_residual=False,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__()
|
||||
ff_kwargs, kwargs = groupby_prefix_and_trim('ff_', kwargs)
|
||||
attn_kwargs, _ = groupby_prefix_and_trim('attn_', kwargs)
|
||||
|
||||
dim_head = attn_kwargs.get('dim_head', DEFAULT_DIM_HEAD)
|
||||
|
||||
self.dim = dim
|
||||
self.depth = depth
|
||||
self.layers = nn.ModuleList([])
|
||||
|
||||
self.has_pos_emb = position_infused_attn
|
||||
self.pia_pos_emb = FixedPositionalEmbedding(dim) if position_infused_attn else None
|
||||
self.rotary_pos_emb = always(None)
|
||||
|
||||
assert rel_pos_num_buckets <= rel_pos_max_distance, 'number of relative position buckets must be less than the relative position max distance'
|
||||
self.rel_pos = None
|
||||
|
||||
self.pre_norm = pre_norm
|
||||
|
||||
self.residual_attn = residual_attn
|
||||
self.cross_residual_attn = cross_residual_attn
|
||||
|
||||
norm_class = ScaleNorm if use_scalenorm else nn.LayerNorm
|
||||
norm_class = RMSNorm if use_rmsnorm else norm_class
|
||||
norm_fn = partial(norm_class, dim)
|
||||
|
||||
norm_fn = nn.Identity if use_rezero else norm_fn
|
||||
branch_fn = Rezero if use_rezero else None
|
||||
|
||||
if cross_attend and not only_cross:
|
||||
default_block = ('a', 'c', 'f')
|
||||
elif cross_attend and only_cross:
|
||||
default_block = ('c', 'f')
|
||||
else:
|
||||
default_block = ('a', 'f')
|
||||
|
||||
if macaron:
|
||||
default_block = ('f',) + default_block
|
||||
|
||||
if exists(custom_layers):
|
||||
layer_types = custom_layers
|
||||
elif exists(par_ratio):
|
||||
par_depth = depth * len(default_block)
|
||||
assert 1 < par_ratio <= par_depth, 'par ratio out of range'
|
||||
default_block = tuple(filter(not_equals('f'), default_block))
|
||||
par_attn = par_depth // par_ratio
|
||||
depth_cut = par_depth * 2 // 3 # 2 / 3 attention layer cutoff suggested by PAR paper
|
||||
par_width = (depth_cut + depth_cut // par_attn) // par_attn
|
||||
assert len(default_block) <= par_width, 'default block is too large for par_ratio'
|
||||
par_block = default_block + ('f',) * (par_width - len(default_block))
|
||||
par_head = par_block * par_attn
|
||||
layer_types = par_head + ('f',) * (par_depth - len(par_head))
|
||||
elif exists(sandwich_coef):
|
||||
assert sandwich_coef > 0 and sandwich_coef <= depth, 'sandwich coefficient should be less than the depth'
|
||||
layer_types = ('a',) * sandwich_coef + default_block * (depth - sandwich_coef) + ('f',) * sandwich_coef
|
||||
else:
|
||||
layer_types = default_block * depth
|
||||
|
||||
self.layer_types = layer_types
|
||||
self.num_attn_layers = len(list(filter(equals('a'), layer_types)))
|
||||
|
||||
for layer_type in self.layer_types:
|
||||
if layer_type == 'a':
|
||||
layer = Attention(dim, heads=heads, causal=causal, **attn_kwargs)
|
||||
elif layer_type == 'c':
|
||||
layer = Attention(dim, heads=heads, **attn_kwargs)
|
||||
elif layer_type == 'f':
|
||||
layer = FeedForward(dim, **ff_kwargs)
|
||||
layer = layer if not macaron else Scale(0.5, layer)
|
||||
else:
|
||||
raise Exception(f'invalid layer type {layer_type}')
|
||||
|
||||
if isinstance(layer, Attention) and exists(branch_fn):
|
||||
layer = branch_fn(layer)
|
||||
|
||||
if gate_residual:
|
||||
residual_fn = GRUGating(dim)
|
||||
else:
|
||||
residual_fn = Residual()
|
||||
|
||||
self.layers.append(nn.ModuleList([
|
||||
norm_fn(),
|
||||
layer,
|
||||
residual_fn
|
||||
]))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
context=None,
|
||||
mask=None,
|
||||
context_mask=None,
|
||||
mems=None,
|
||||
return_hiddens=False
|
||||
):
|
||||
hiddens = []
|
||||
intermediates = []
|
||||
prev_attn = None
|
||||
prev_cross_attn = None
|
||||
|
||||
mems = mems.copy() if exists(mems) else [None] * self.num_attn_layers
|
||||
|
||||
for ind, (layer_type, (norm, block, residual_fn)) in enumerate(zip(self.layer_types, self.layers)):
|
||||
is_last = ind == (len(self.layers) - 1)
|
||||
|
||||
if layer_type == 'a':
|
||||
hiddens.append(x)
|
||||
layer_mem = mems.pop(0)
|
||||
|
||||
residual = x
|
||||
|
||||
if self.pre_norm:
|
||||
x = norm(x)
|
||||
|
||||
if layer_type == 'a':
|
||||
out, inter = block(x, mask=mask, sinusoidal_emb=self.pia_pos_emb, rel_pos=self.rel_pos,
|
||||
prev_attn=prev_attn, mem=layer_mem)
|
||||
elif layer_type == 'c':
|
||||
out, inter = block(x, context=context, mask=mask, context_mask=context_mask, prev_attn=prev_cross_attn)
|
||||
elif layer_type == 'f':
|
||||
out = block(x)
|
||||
|
||||
x = residual_fn(out, residual)
|
||||
|
||||
if layer_type in ('a', 'c'):
|
||||
intermediates.append(inter)
|
||||
|
||||
if layer_type == 'a' and self.residual_attn:
|
||||
prev_attn = inter.pre_softmax_attn
|
||||
elif layer_type == 'c' and self.cross_residual_attn:
|
||||
prev_cross_attn = inter.pre_softmax_attn
|
||||
|
||||
if not self.pre_norm and not is_last:
|
||||
x = norm(x)
|
||||
|
||||
if return_hiddens:
|
||||
intermediates = LayerIntermediates(
|
||||
hiddens=hiddens,
|
||||
attn_intermediates=intermediates
|
||||
)
|
||||
|
||||
return x, intermediates
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class Encoder(AttentionLayers):
|
||||
def __init__(self, **kwargs):
|
||||
assert 'causal' not in kwargs, 'cannot set causality on encoder'
|
||||
super().__init__(causal=False, **kwargs)
|
||||
|
||||
|
||||
|
||||
class TransformerWrapper(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
num_tokens,
|
||||
max_seq_len,
|
||||
attn_layers,
|
||||
emb_dim=None,
|
||||
max_mem_len=0.,
|
||||
emb_dropout=0.,
|
||||
num_memory_tokens=None,
|
||||
tie_embedding=False,
|
||||
use_pos_emb=True
|
||||
):
|
||||
super().__init__()
|
||||
assert isinstance(attn_layers, AttentionLayers), 'attention layers must be one of Encoder or Decoder'
|
||||
|
||||
dim = attn_layers.dim
|
||||
emb_dim = default(emb_dim, dim)
|
||||
|
||||
self.max_seq_len = max_seq_len
|
||||
self.max_mem_len = max_mem_len
|
||||
self.num_tokens = num_tokens
|
||||
|
||||
self.token_emb = nn.Embedding(num_tokens, emb_dim)
|
||||
self.pos_emb = AbsolutePositionalEmbedding(emb_dim, max_seq_len) if (
|
||||
use_pos_emb and not attn_layers.has_pos_emb) else always(0)
|
||||
self.emb_dropout = nn.Dropout(emb_dropout)
|
||||
|
||||
self.project_emb = nn.Linear(emb_dim, dim) if emb_dim != dim else nn.Identity()
|
||||
self.attn_layers = attn_layers
|
||||
self.norm = nn.LayerNorm(dim)
|
||||
|
||||
self.init_()
|
||||
|
||||
self.to_logits = nn.Linear(dim, num_tokens) if not tie_embedding else lambda t: t @ self.token_emb.weight.t()
|
||||
|
||||
# memory tokens (like [cls]) from Memory Transformers paper
|
||||
num_memory_tokens = default(num_memory_tokens, 0)
|
||||
self.num_memory_tokens = num_memory_tokens
|
||||
if num_memory_tokens > 0:
|
||||
self.memory_tokens = nn.Parameter(torch.randn(num_memory_tokens, dim))
|
||||
|
||||
# let funnel encoder know number of memory tokens, if specified
|
||||
if hasattr(attn_layers, 'num_memory_tokens'):
|
||||
attn_layers.num_memory_tokens = num_memory_tokens
|
||||
|
||||
def init_(self):
|
||||
nn.init.normal_(self.token_emb.weight, std=0.02)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
return_embeddings=False,
|
||||
mask=None,
|
||||
return_mems=False,
|
||||
return_attn=False,
|
||||
mems=None,
|
||||
**kwargs
|
||||
):
|
||||
b, n, device, num_mem = *x.shape, x.device, self.num_memory_tokens
|
||||
x = self.token_emb(x)
|
||||
x += self.pos_emb(x)
|
||||
x = self.emb_dropout(x)
|
||||
|
||||
x = self.project_emb(x)
|
||||
|
||||
if num_mem > 0:
|
||||
mem = repeat(self.memory_tokens, 'n d -> b n d', b=b)
|
||||
x = torch.cat((mem, x), dim=1)
|
||||
|
||||
# auto-handle masking after appending memory tokens
|
||||
if exists(mask):
|
||||
mask = F.pad(mask, (num_mem, 0), value=True)
|
||||
|
||||
x, intermediates = self.attn_layers(x, mask=mask, mems=mems, return_hiddens=True, **kwargs)
|
||||
x = self.norm(x)
|
||||
|
||||
mem, x = x[:, :num_mem], x[:, num_mem:]
|
||||
|
||||
out = self.to_logits(x) if not return_embeddings else x
|
||||
|
||||
if return_mems:
|
||||
hiddens = intermediates.hiddens
|
||||
new_mems = list(map(lambda pair: torch.cat(pair, dim=-2), zip(mems, hiddens))) if exists(mems) else hiddens
|
||||
new_mems = list(map(lambda t: t[..., -self.max_mem_len:, :].detach(), new_mems))
|
||||
return out, new_mems
|
||||
|
||||
if return_attn:
|
||||
attn_maps = list(map(lambda t: t.post_softmax_attn, intermediates.attn_intermediates))
|
||||
return out, attn_maps
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,143 @@
|
||||
|
||||
import torch
|
||||
|
||||
from collections import OrderedDict
|
||||
|
||||
from comfy import model_base
|
||||
from comfy import utils
|
||||
from comfy import diffusers_convert
|
||||
|
||||
from comfy import sd2_clip
|
||||
|
||||
from comfy import supported_models_base
|
||||
from comfy import latent_formats
|
||||
|
||||
from ..lvdm.modules.encoders.resampler import Resampler
|
||||
|
||||
DYNAMICRAFTER_CONFIG = {
|
||||
'in_channels': 8,
|
||||
'out_channels': 4,
|
||||
'model_channels': 320,
|
||||
'attention_resolutions': [4, 2, 1],
|
||||
'num_res_blocks': 2,
|
||||
'channel_mult': [1, 2, 4, 4],
|
||||
'num_head_channels': 64,
|
||||
'transformer_depth': 1,
|
||||
'context_dim': 1024,
|
||||
'use_linear': True,
|
||||
'use_checkpoint': False,
|
||||
'temporal_conv': True,
|
||||
'temporal_attention': True,
|
||||
'temporal_selfatt_only': True,
|
||||
'use_relative_position': False,
|
||||
'use_causal_attention': False,
|
||||
'temporal_length': 16,
|
||||
'addition_attention': True,
|
||||
'image_cross_attention': True,
|
||||
'image_cross_attention_scale_learnable': True,
|
||||
'default_fs': 3,
|
||||
'fs_condition': True
|
||||
}
|
||||
|
||||
IMAGE_PROJ_CONFIG = {
|
||||
"dim": 1024,
|
||||
"depth": 4,
|
||||
"dim_head": 64,
|
||||
"heads": 12,
|
||||
"num_queries": 16,
|
||||
"embedding_dim": 1280,
|
||||
"output_dim": 1024,
|
||||
"ff_mult": 4,
|
||||
"video_length": 16
|
||||
}
|
||||
|
||||
def process_list_or_str(target_key_or_keys, k):
|
||||
if isinstance(target_key_or_keys, list):
|
||||
return any([list_k in k for list_k in target_key_or_keys])
|
||||
else:
|
||||
return target_key_or_keys in k
|
||||
|
||||
def simple_state_dict_loader(state_dict: dict, target_key: str, target_dict: dict = None):
|
||||
out_dict = {}
|
||||
|
||||
if target_dict is None:
|
||||
for k, v in state_dict.items():
|
||||
if process_list_or_str(target_key, k):
|
||||
out_dict[k] = v
|
||||
else:
|
||||
for k, v in target_dict.items():
|
||||
out_dict[k] = state_dict[k]
|
||||
|
||||
return out_dict
|
||||
|
||||
def load_image_proj_dict(state_dict: dict):
|
||||
return simple_state_dict_loader(state_dict, 'image_proj')
|
||||
|
||||
def load_dynamicrafter_dict(state_dict: dict):
|
||||
return simple_state_dict_loader(state_dict, 'model.diffusion_model')
|
||||
|
||||
def load_vae_dict(state_dict: dict):
|
||||
return simple_state_dict_loader(state_dict, 'first_stage_model')
|
||||
|
||||
def get_base_model(state_dict: dict, version_checker=False):
|
||||
|
||||
is_256_model = False
|
||||
|
||||
for k in state_dict.keys():
|
||||
if "framestride_embed" in k:
|
||||
is_256_model = True
|
||||
break
|
||||
|
||||
def get_image_proj_model(state_dict: dict):
|
||||
|
||||
state_dict = {k.replace('image_proj_model.', ''): v for k, v in state_dict.items()}
|
||||
#target_dict = Resampler().state_dict()
|
||||
|
||||
ImageProjModel = Resampler(**IMAGE_PROJ_CONFIG)
|
||||
ImageProjModel.load_state_dict(state_dict)
|
||||
|
||||
print("Image Projection Model loaded successfully")
|
||||
#del target_dict
|
||||
return ImageProjModel
|
||||
|
||||
class DynamiCrafterBase(supported_models_base.BASE):
|
||||
unet_config = {}
|
||||
unet_extra_config = {}
|
||||
|
||||
latent_format = latent_formats.SD15
|
||||
|
||||
def process_clip_state_dict(self, state_dict):
|
||||
replace_prefix = {}
|
||||
replace_prefix["conditioner.embedders.0.model."] = "clip_h." #SD2 in sgm format
|
||||
replace_prefix["cond_stage_model.model."] = "clip_h."
|
||||
state_dict = utils.state_dict_prefix_replace(state_dict, replace_prefix, filter_keys=True)
|
||||
state_dict = utils.clip_text_transformers_convert(state_dict, "clip_h.", "clip_h.transformer.")
|
||||
return state_dict
|
||||
|
||||
def process_clip_state_dict_for_saving(self, state_dict):
|
||||
replace_prefix = {}
|
||||
replace_prefix["clip_h"] = "cond_stage_model.model"
|
||||
state_dict = utils.state_dict_prefix_replace(state_dict, replace_prefix)
|
||||
state_dict = diffusers_convert.convert_text_enc_state_dict_v20(state_dict)
|
||||
return state_dict
|
||||
|
||||
def clip_target(self):
|
||||
return supported_models_base.ClipTarget(sd2_clip.SD2Tokenizer, sd2_clip.SD2ClipModel)
|
||||
|
||||
def process_dict_version(self, state_dict: dict):
|
||||
processed_dict = OrderedDict()
|
||||
is_eps = False
|
||||
|
||||
for k in list(state_dict.keys()):
|
||||
if "framestride_embed" in k:
|
||||
new_key = k.replace("framestride_embed", "fps_embedding")
|
||||
processed_dict[new_key] = state_dict[k]
|
||||
is_eps = True
|
||||
continue
|
||||
|
||||
processed_dict[k] = state_dict[k]
|
||||
|
||||
return processed_dict, is_eps
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,82 @@
|
||||
import importlib
|
||||
import numpy as np
|
||||
import cv2
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
MODEL_EXTS = ['ckpt', 'safetensors', 'bin']
|
||||
|
||||
def get_models_directory(directory: list):
|
||||
files_list = list(filter(lambda f: f.split(".")[-1] in MODEL_EXTS, directory))
|
||||
return files_list
|
||||
|
||||
def count_params(model, verbose=False):
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
if verbose:
|
||||
print(f"{model.__class__.__name__} has {total_params*1.e-6:.2f} M params.")
|
||||
return total_params
|
||||
|
||||
|
||||
def check_istarget(name, para_list):
|
||||
"""
|
||||
name: full name of source para
|
||||
para_list: partial name of target para
|
||||
"""
|
||||
istarget=False
|
||||
for para in para_list:
|
||||
if para in name:
|
||||
return True
|
||||
return istarget
|
||||
|
||||
|
||||
def instantiate_from_config(config):
|
||||
if not "target" in config:
|
||||
if config == '__is_first_stage__':
|
||||
return None
|
||||
elif config == "__is_unconditional__":
|
||||
return None
|
||||
raise KeyError("Expected key `target` to instantiate.")
|
||||
return get_obj_from_str(config["target"])(**config.get("params", dict()))
|
||||
|
||||
|
||||
def get_obj_from_str(string, reload=False):
|
||||
module, cls = string.rsplit(".", 1)
|
||||
if reload:
|
||||
module_imp = importlib.import_module(module)
|
||||
importlib.reload(module_imp)
|
||||
return getattr(importlib.import_module(module, package=None), cls)
|
||||
|
||||
|
||||
def load_npz_from_dir(data_dir):
|
||||
data = [np.load(os.path.join(data_dir, data_name))['arr_0'] for data_name in os.listdir(data_dir)]
|
||||
data = np.concatenate(data, axis=0)
|
||||
return data
|
||||
|
||||
|
||||
def load_npz_from_paths(data_paths):
|
||||
data = [np.load(data_path)['arr_0'] for data_path in data_paths]
|
||||
data = np.concatenate(data, axis=0)
|
||||
return data
|
||||
|
||||
|
||||
def resize_numpy_image(image, max_resolution=512 * 512, resize_short_edge=None):
|
||||
h, w = image.shape[:2]
|
||||
if resize_short_edge is not None:
|
||||
k = resize_short_edge / min(h, w)
|
||||
else:
|
||||
k = max_resolution / (h * w)
|
||||
k = k**0.5
|
||||
h = int(np.round(h * k / 64)) * 64
|
||||
w = int(np.round(w * k / 64)) * 64
|
||||
image = cv2.resize(image, (w, h), interpolation=cv2.INTER_LANCZOS4)
|
||||
return image
|
||||
|
||||
|
||||
def setup_dist(args):
|
||||
if dist.is_initialized():
|
||||
return
|
||||
torch.cuda.set_device(args.local_rank)
|
||||
torch.distributed.init_process_group(
|
||||
'nccl',
|
||||
init_method='env://'
|
||||
)
|
||||
+2699
-573
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,156 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
import cv2
|
||||
import torchvision.transforms as transforms
|
||||
from torch.utils.data import DataLoader
|
||||
from .simple_extractor_dataset import SimpleFolderDataset
|
||||
from .transforms import transform_logits
|
||||
from tqdm import tqdm
|
||||
from PIL import Image
|
||||
|
||||
def get_palette(num_cls):
|
||||
""" Returns the color map for visualizing the segmentation mask.
|
||||
Args:
|
||||
num_cls: Number of classes
|
||||
Returns:
|
||||
The color map
|
||||
"""
|
||||
n = num_cls
|
||||
palette = [0] * (n * 3)
|
||||
for j in range(0, n):
|
||||
lab = j
|
||||
palette[j * 3 + 0] = 0
|
||||
palette[j * 3 + 1] = 0
|
||||
palette[j * 3 + 2] = 0
|
||||
i = 0
|
||||
while lab:
|
||||
palette[j * 3 + 0] |= (((lab >> 0) & 1) << (7 - i))
|
||||
palette[j * 3 + 1] |= (((lab >> 1) & 1) << (7 - i))
|
||||
palette[j * 3 + 2] |= (((lab >> 2) & 1) << (7 - i))
|
||||
i += 1
|
||||
lab >>= 3
|
||||
return palette
|
||||
|
||||
|
||||
def delete_irregular(logits_result):
|
||||
parsing_result = np.argmax(logits_result, axis=2)
|
||||
upper_cloth = np.where(parsing_result == 4, 255, 0)
|
||||
contours, hierarchy = cv2.findContours(upper_cloth.astype(np.uint8),
|
||||
cv2.RETR_CCOMP, cv2.CHAIN_APPROX_TC89_L1)
|
||||
area = []
|
||||
for i in range(len(contours)):
|
||||
a = cv2.contourArea(contours[i], True)
|
||||
area.append(abs(a))
|
||||
if len(area) != 0:
|
||||
top = area.index(max(area))
|
||||
M = cv2.moments(contours[top])
|
||||
cY = int(M["m01"] / M["m00"])
|
||||
|
||||
dresses = np.where(parsing_result == 7, 255, 0)
|
||||
contours_dress, hierarchy_dress = cv2.findContours(dresses.astype(np.uint8),
|
||||
cv2.RETR_CCOMP, cv2.CHAIN_APPROX_TC89_L1)
|
||||
area_dress = []
|
||||
for j in range(len(contours_dress)):
|
||||
a_d = cv2.contourArea(contours_dress[j], True)
|
||||
area_dress.append(abs(a_d))
|
||||
if len(area_dress) != 0:
|
||||
top_dress = area_dress.index(max(area_dress))
|
||||
M_dress = cv2.moments(contours_dress[top_dress])
|
||||
cY_dress = int(M_dress["m01"] / M_dress["m00"])
|
||||
wear_type = "dresses"
|
||||
if len(area) != 0:
|
||||
if len(area_dress) != 0 and cY_dress > cY:
|
||||
irregular_list = np.array([4, 5, 6])
|
||||
logits_result[:, :, irregular_list] = -1
|
||||
else:
|
||||
irregular_list = np.array([5, 6, 7, 8, 9, 10, 12, 13])
|
||||
logits_result[:cY, :, irregular_list] = -1
|
||||
wear_type = "cloth_pant"
|
||||
parsing_result = np.argmax(logits_result, axis=2)
|
||||
# pad border
|
||||
parsing_result = np.pad(parsing_result, pad_width=1, mode='constant', constant_values=0)
|
||||
return parsing_result, wear_type
|
||||
|
||||
|
||||
|
||||
def hole_fill(img):
|
||||
img_copy = img.copy()
|
||||
mask = np.zeros((img.shape[0] + 2, img.shape[1] + 2), dtype=np.uint8)
|
||||
cv2.floodFill(img, mask, (0, 0), 255)
|
||||
img_inverse = cv2.bitwise_not(img)
|
||||
dst = cv2.bitwise_or(img_copy, img_inverse)
|
||||
return dst
|
||||
|
||||
def refine_mask(mask):
|
||||
contours, hierarchy = cv2.findContours(mask.astype(np.uint8),
|
||||
cv2.RETR_CCOMP, cv2.CHAIN_APPROX_TC89_L1)
|
||||
area = []
|
||||
for j in range(len(contours)):
|
||||
a_d = cv2.contourArea(contours[j], True)
|
||||
area.append(abs(a_d))
|
||||
refine_mask = np.zeros_like(mask).astype(np.uint8)
|
||||
if len(area) != 0:
|
||||
i = area.index(max(area))
|
||||
cv2.drawContours(refine_mask, contours, i, color=255, thickness=-1)
|
||||
# keep large area in skin case
|
||||
for j in range(len(area)):
|
||||
if j != i and area[i] > 2000:
|
||||
cv2.drawContours(refine_mask, contours, j, color=255, thickness=-1)
|
||||
return refine_mask
|
||||
|
||||
def refine_hole(parsing_result_filled, parsing_result, arm_mask):
|
||||
filled_hole = cv2.bitwise_and(np.where(parsing_result_filled == 4, 255, 0),
|
||||
np.where(parsing_result != 4, 255, 0)) - arm_mask * 255
|
||||
contours, hierarchy = cv2.findContours(filled_hole, cv2.RETR_CCOMP, cv2.CHAIN_APPROX_TC89_L1)
|
||||
refine_hole_mask = np.zeros_like(parsing_result).astype(np.uint8)
|
||||
for i in range(len(contours)):
|
||||
a = cv2.contourArea(contours[i], True)
|
||||
# keep hole > 2000 pixels
|
||||
if abs(a) > 2000:
|
||||
cv2.drawContours(refine_hole_mask, contours, i, color=255, thickness=-1)
|
||||
return refine_hole_mask + arm_mask
|
||||
|
||||
def onnx_inference(lip_session, input_dir, mask_components=[0]):
|
||||
|
||||
transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=[0.406, 0.456, 0.485], std=[0.225, 0.224, 0.229])
|
||||
])
|
||||
input_size = [473, 473]
|
||||
|
||||
dataset_lip = SimpleFolderDataset(root=input_dir, input_size=input_size, transform=transform)
|
||||
dataloader_lip = DataLoader(dataset_lip)
|
||||
palette = get_palette(20)
|
||||
with torch.no_grad():
|
||||
for _, batch in enumerate(tqdm(dataloader_lip)):
|
||||
image, meta = batch
|
||||
c = meta['center'].numpy()[0]
|
||||
s = meta['scale'].numpy()[0]
|
||||
w = meta['width'].numpy()[0]
|
||||
h = meta['height'].numpy()[0]
|
||||
|
||||
output = lip_session.run(None, {"input.1": image.numpy().astype(np.float32)})
|
||||
upsample = torch.nn.Upsample(size=input_size, mode='bilinear', align_corners=True)
|
||||
upsample_output = upsample(torch.from_numpy(output[1][0]).unsqueeze(0))
|
||||
upsample_output = upsample_output.squeeze()
|
||||
upsample_output = upsample_output.permute(1, 2, 0) # CHW -> HWC
|
||||
logits_result_lip = transform_logits(upsample_output.data.cpu().numpy(), c, s, w, h,
|
||||
input_size=input_size)
|
||||
parsing_result = np.argmax(logits_result_lip, axis=2)
|
||||
|
||||
output_img = Image.fromarray(np.asarray(parsing_result, dtype=np.uint8))
|
||||
output_img.putpalette(palette)
|
||||
|
||||
mask = np.isin(output_img, mask_components).astype(np.uint8)
|
||||
mask_image = Image.fromarray(mask * 255)
|
||||
mask_image = mask_image.convert("RGB")
|
||||
mask_image = torch.from_numpy(np.array(mask_image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
output_img = output_img.convert('RGB')
|
||||
output_img = torch.from_numpy(np.array(output_img).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
return output_img, mask_image
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
from .parsing_api import onnx_inference
|
||||
from ..libs.utils import install_package
|
||||
|
||||
class HumanParsing:
|
||||
def __init__(self, model_path):
|
||||
self.model_path = model_path
|
||||
self.session = None
|
||||
|
||||
def __call__(self, input_image, mask_components):
|
||||
if self.session is None:
|
||||
install_package('onnxruntime')
|
||||
import onnxruntime as ort
|
||||
|
||||
session_options = ort.SessionOptions()
|
||||
session_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
|
||||
session_options.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL
|
||||
# session_options.add_session_config_entry('gpu_id', str(gpu_id))
|
||||
self.session = ort.InferenceSession(self.model_path, sess_options=session_options,
|
||||
providers=['CPUExecutionProvider'])
|
||||
|
||||
parsed_image, mask = onnx_inference(self.session, input_image, mask_components)
|
||||
return parsed_image, mask
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
#!/usr/bin/env python
|
||||
# -*- encoding: utf-8 -*-
|
||||
|
||||
"""
|
||||
@Author : Peike Li
|
||||
@Contact : peike.li@yahoo.com
|
||||
@File : dataset.py
|
||||
@Time : 8/30/19 9:12 PM
|
||||
@Desc : Dataset Definition
|
||||
@License : This source code is licensed under the license found in the
|
||||
LICENSE file in the root directory of this source tree.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from torch.utils import data
|
||||
from .transforms import get_affine_transform
|
||||
|
||||
|
||||
class SimpleFolderDataset(data.Dataset):
|
||||
def __init__(self, root, input_size=[512, 512], transform=None):
|
||||
self.root = root
|
||||
self.input_size = input_size
|
||||
self.transform = transform
|
||||
self.aspect_ratio = input_size[1] * 1.0 / input_size[0]
|
||||
self.input_size = np.asarray(input_size)
|
||||
self.is_pil_image = False
|
||||
if isinstance(root, Image.Image):
|
||||
self.file_list = [root]
|
||||
self.is_pil_image = True
|
||||
elif os.path.isfile(root):
|
||||
self.file_list = [os.path.basename(root)]
|
||||
self.root = os.path.dirname(root)
|
||||
else:
|
||||
self.file_list = os.listdir(self.root)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.file_list)
|
||||
|
||||
def _box2cs(self, box):
|
||||
x, y, w, h = box[:4]
|
||||
return self._xywh2cs(x, y, w, h)
|
||||
|
||||
def _xywh2cs(self, x, y, w, h):
|
||||
center = np.zeros((2), dtype=np.float32)
|
||||
center[0] = x + w * 0.5
|
||||
center[1] = y + h * 0.5
|
||||
if w > self.aspect_ratio * h:
|
||||
h = w * 1.0 / self.aspect_ratio
|
||||
elif w < self.aspect_ratio * h:
|
||||
w = h * self.aspect_ratio
|
||||
scale = np.array([w, h], dtype=np.float32)
|
||||
return center, scale
|
||||
|
||||
def __getitem__(self, index):
|
||||
if self.is_pil_image:
|
||||
img = np.asarray(self.file_list[index])[:, :, [2, 1, 0]]
|
||||
else:
|
||||
img_name = self.file_list[index]
|
||||
img_path = os.path.join(self.root, img_name)
|
||||
img = cv2.imread(img_path, cv2.IMREAD_COLOR)
|
||||
h, w, _ = img.shape
|
||||
|
||||
# Get person center and scale
|
||||
person_center, s = self._box2cs([0, 0, w - 1, h - 1])
|
||||
r = 0
|
||||
trans = get_affine_transform(person_center, s, r, self.input_size)
|
||||
input = cv2.warpAffine(
|
||||
img,
|
||||
trans,
|
||||
(int(self.input_size[1]), int(self.input_size[0])),
|
||||
flags=cv2.INTER_LINEAR,
|
||||
borderMode=cv2.BORDER_CONSTANT,
|
||||
borderValue=(0, 0, 0))
|
||||
|
||||
input = self.transform(input)
|
||||
meta = {
|
||||
'center': person_center,
|
||||
'height': h,
|
||||
'width': w,
|
||||
'scale': s,
|
||||
'rotation': r
|
||||
}
|
||||
|
||||
return input, meta
|
||||
@@ -0,0 +1,167 @@
|
||||
# ------------------------------------------------------------------------------
|
||||
# Copyright (c) Microsoft
|
||||
# Licensed under the MIT License.
|
||||
# Written by Bin Xiao (Bin.Xiao@microsoft.com)
|
||||
# ------------------------------------------------------------------------------
|
||||
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
|
||||
import numpy as np
|
||||
import cv2
|
||||
import torch
|
||||
|
||||
class BRG2Tensor_transform(object):
|
||||
def __call__(self, pic):
|
||||
img = torch.from_numpy(pic.transpose((2, 0, 1)))
|
||||
if isinstance(img, torch.ByteTensor):
|
||||
return img.float()
|
||||
else:
|
||||
return img
|
||||
|
||||
class BGR2RGB_transform(object):
|
||||
def __call__(self, tensor):
|
||||
return tensor[[2,1,0],:,:]
|
||||
|
||||
def flip_back(output_flipped, matched_parts):
|
||||
'''
|
||||
ouput_flipped: numpy.ndarray(batch_size, num_joints, height, width)
|
||||
'''
|
||||
assert output_flipped.ndim == 4,\
|
||||
'output_flipped should be [batch_size, num_joints, height, width]'
|
||||
|
||||
output_flipped = output_flipped[:, :, :, ::-1]
|
||||
|
||||
for pair in matched_parts:
|
||||
tmp = output_flipped[:, pair[0], :, :].copy()
|
||||
output_flipped[:, pair[0], :, :] = output_flipped[:, pair[1], :, :]
|
||||
output_flipped[:, pair[1], :, :] = tmp
|
||||
|
||||
return output_flipped
|
||||
|
||||
|
||||
def fliplr_joints(joints, joints_vis, width, matched_parts):
|
||||
"""
|
||||
flip coords
|
||||
"""
|
||||
# Flip horizontal
|
||||
joints[:, 0] = width - joints[:, 0] - 1
|
||||
|
||||
# Change left-right parts
|
||||
for pair in matched_parts:
|
||||
joints[pair[0], :], joints[pair[1], :] = \
|
||||
joints[pair[1], :], joints[pair[0], :].copy()
|
||||
joints_vis[pair[0], :], joints_vis[pair[1], :] = \
|
||||
joints_vis[pair[1], :], joints_vis[pair[0], :].copy()
|
||||
|
||||
return joints*joints_vis, joints_vis
|
||||
|
||||
|
||||
def transform_preds(coords, center, scale, input_size):
|
||||
target_coords = np.zeros(coords.shape)
|
||||
trans = get_affine_transform(center, scale, 0, input_size, inv=1)
|
||||
for p in range(coords.shape[0]):
|
||||
target_coords[p, 0:2] = affine_transform(coords[p, 0:2], trans)
|
||||
return target_coords
|
||||
|
||||
def transform_parsing(pred, center, scale, width, height, input_size):
|
||||
|
||||
trans = get_affine_transform(center, scale, 0, input_size, inv=1)
|
||||
target_pred = cv2.warpAffine(
|
||||
pred,
|
||||
trans,
|
||||
(int(width), int(height)), #(int(width), int(height)),
|
||||
flags=cv2.INTER_NEAREST,
|
||||
borderMode=cv2.BORDER_CONSTANT,
|
||||
borderValue=(0))
|
||||
|
||||
return target_pred
|
||||
|
||||
def transform_logits(logits, center, scale, width, height, input_size):
|
||||
|
||||
trans = get_affine_transform(center, scale, 0, input_size, inv=1)
|
||||
channel = logits.shape[2]
|
||||
target_logits = []
|
||||
for i in range(channel):
|
||||
target_logit = cv2.warpAffine(
|
||||
logits[:,:,i],
|
||||
trans,
|
||||
(int(width), int(height)), #(int(width), int(height)),
|
||||
flags=cv2.INTER_LINEAR,
|
||||
borderMode=cv2.BORDER_CONSTANT,
|
||||
borderValue=(0))
|
||||
target_logits.append(target_logit)
|
||||
target_logits = np.stack(target_logits,axis=2)
|
||||
|
||||
return target_logits
|
||||
|
||||
|
||||
def get_affine_transform(center,
|
||||
scale,
|
||||
rot,
|
||||
output_size,
|
||||
shift=np.array([0, 0], dtype=np.float32),
|
||||
inv=0):
|
||||
if not isinstance(scale, np.ndarray) and not isinstance(scale, list):
|
||||
print(scale)
|
||||
scale = np.array([scale, scale])
|
||||
|
||||
scale_tmp = scale
|
||||
|
||||
src_w = scale_tmp[0]
|
||||
dst_w = output_size[1]
|
||||
dst_h = output_size[0]
|
||||
|
||||
rot_rad = np.pi * rot / 180
|
||||
src_dir = get_dir([0, src_w * -0.5], rot_rad)
|
||||
dst_dir = np.array([0, (dst_w-1) * -0.5], np.float32)
|
||||
|
||||
src = np.zeros((3, 2), dtype=np.float32)
|
||||
dst = np.zeros((3, 2), dtype=np.float32)
|
||||
src[0, :] = center + scale_tmp * shift
|
||||
src[1, :] = center + src_dir + scale_tmp * shift
|
||||
dst[0, :] = [(dst_w-1) * 0.5, (dst_h-1) * 0.5]
|
||||
dst[1, :] = np.array([(dst_w-1) * 0.5, (dst_h-1) * 0.5]) + dst_dir
|
||||
|
||||
src[2:, :] = get_3rd_point(src[0, :], src[1, :])
|
||||
dst[2:, :] = get_3rd_point(dst[0, :], dst[1, :])
|
||||
|
||||
if inv:
|
||||
trans = cv2.getAffineTransform(np.float32(dst), np.float32(src))
|
||||
else:
|
||||
trans = cv2.getAffineTransform(np.float32(src), np.float32(dst))
|
||||
|
||||
return trans
|
||||
|
||||
|
||||
def affine_transform(pt, t):
|
||||
new_pt = np.array([pt[0], pt[1], 1.]).T
|
||||
new_pt = np.dot(t, new_pt)
|
||||
return new_pt[:2]
|
||||
|
||||
|
||||
def get_3rd_point(a, b):
|
||||
direct = a - b
|
||||
return b + np.array([-direct[1], direct[0]], dtype=np.float32)
|
||||
|
||||
|
||||
def get_dir(src_point, rot_rad):
|
||||
sn, cs = np.sin(rot_rad), np.cos(rot_rad)
|
||||
|
||||
src_result = [0, 0]
|
||||
src_result[0] = src_point[0] * cs - src_point[1] * sn
|
||||
src_result[1] = src_point[0] * sn + src_point[1] * cs
|
||||
|
||||
return src_result
|
||||
|
||||
|
||||
def crop(img, center, scale, output_size, rot=0):
|
||||
trans = get_affine_transform(center, scale, rot, output_size)
|
||||
|
||||
dst_img = cv2.warpAffine(img,
|
||||
trans,
|
||||
(int(output_size[1]), int(output_size[0])),
|
||||
flags=cv2.INTER_LINEAR)
|
||||
|
||||
return dst_img
|
||||
@@ -0,0 +1,185 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
from typing import Tuple, TypedDict, Callable
|
||||
|
||||
import comfy.model_management
|
||||
from comfy.sd import load_unet
|
||||
from comfy.ldm.models.autoencoder import AutoencoderKL
|
||||
from comfy.model_base import BaseModel
|
||||
from PIL import Image
|
||||
from nodes import VAEEncode
|
||||
|
||||
from ..layer_diffuse.model import ModelPatcher, calculate_weight_adjust_channel
|
||||
from ..libs.image import np2tensor, pil2tensor
|
||||
|
||||
class UnetParams(TypedDict):
|
||||
input: torch.Tensor
|
||||
timestep: torch.Tensor
|
||||
c: dict
|
||||
cond_or_uncond: torch.Tensor
|
||||
|
||||
|
||||
class VAEEncodeArgMax(VAEEncode):
|
||||
def encode(self, vae, pixels):
|
||||
assert isinstance(
|
||||
vae.first_stage_model, AutoencoderKL
|
||||
), "ArgMax only supported for AutoencoderKL"
|
||||
original_sample_mode = vae.first_stage_model.regularization.sample
|
||||
vae.first_stage_model.regularization.sample = False
|
||||
ret = super().encode(vae, pixels)
|
||||
vae.first_stage_model.regularization.sample = original_sample_mode
|
||||
return ret
|
||||
|
||||
class ICLight:
|
||||
|
||||
@staticmethod
|
||||
def apply_c_concat(params: UnetParams, concat_conds) -> UnetParams:
|
||||
"""Apply c_concat on unet call."""
|
||||
sample = params["input"]
|
||||
params["c"]["c_concat"] = torch.cat(
|
||||
(
|
||||
[concat_conds.to(sample.device)]
|
||||
* (sample.shape[0] // concat_conds.shape[0])
|
||||
),
|
||||
dim=0,
|
||||
)
|
||||
return params
|
||||
|
||||
@staticmethod
|
||||
def create_custom_conv(
|
||||
original_conv: torch.nn.Module,
|
||||
dtype: torch.dtype,
|
||||
device=torch.device,
|
||||
) -> torch.nn.Module:
|
||||
with torch.no_grad():
|
||||
new_conv_in = torch.nn.Conv2d(
|
||||
8,
|
||||
original_conv.out_channels,
|
||||
original_conv.kernel_size,
|
||||
original_conv.stride,
|
||||
original_conv.padding,
|
||||
)
|
||||
new_conv_in.weight.zero_()
|
||||
new_conv_in.weight[:, :4, :, :].copy_(original_conv.weight)
|
||||
new_conv_in.bias = original_conv.bias
|
||||
return new_conv_in.to(dtype=dtype, device=device)
|
||||
|
||||
def generate_lighting_image(self, original_image, direction):
|
||||
_, image_height, image_width, _ = original_image.shape
|
||||
match direction:
|
||||
case 'Left Light':
|
||||
gradient = np.linspace(255, 0, image_width)
|
||||
image = np.tile(gradient, (image_height, 1))
|
||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
||||
return np2tensor(input_bg)
|
||||
case 'Right Light':
|
||||
gradient = np.linspace(0, 255, image_width)
|
||||
image = np.tile(gradient, (image_height, 1))
|
||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
||||
return np2tensor(input_bg)
|
||||
case 'Top Light':
|
||||
gradient = np.linspace(255, 0, image_height)[:, None]
|
||||
image = np.tile(gradient, (1, image_width))
|
||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
||||
return np2tensor(input_bg)
|
||||
case 'Bottom Light':
|
||||
gradient = np.linspace(0, 255, image_height)[:, None]
|
||||
image = np.tile(gradient, (1, image_width))
|
||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
||||
return np2tensor(input_bg)
|
||||
case 'Circle Light':
|
||||
x = np.linspace(-1, 1, image_width)
|
||||
y = np.linspace(-1, 1, image_height)
|
||||
x, y = np.meshgrid(x, y)
|
||||
r = np.sqrt(x ** 2 + y ** 2)
|
||||
r = r / r.max()
|
||||
color1 = np.array([0, 0, 0])[np.newaxis, np.newaxis, :]
|
||||
color2 = np.array([255, 255, 255])[np.newaxis, np.newaxis, :]
|
||||
gradient = (color1 * r[..., np.newaxis] + color2 * (1 - r)[..., np.newaxis]).astype(np.uint8)
|
||||
image = pil2tensor(Image.fromarray(gradient))
|
||||
return image
|
||||
case _:
|
||||
image = pil2tensor(Image.new('RGB', (1, 1), (0, 0, 0)))
|
||||
return image
|
||||
|
||||
def generate_source_image(self, original_image, source):
|
||||
batch_size, image_height, image_width, _ = original_image.shape
|
||||
match source:
|
||||
case 'Use Flipped Background Image':
|
||||
if batch_size < 2:
|
||||
raise ValueError('Must be at least 2 image to use flipped background image.')
|
||||
original_image = [img.unsqueeze(0) for img in original_image]
|
||||
image = torch.flip(original_image[1], [2])
|
||||
return image
|
||||
case 'Ambient':
|
||||
input_bg = np.zeros(shape=(image_height, image_width, 3), dtype=np.uint8) + 64
|
||||
return np2tensor(input_bg)
|
||||
case 'Left Light':
|
||||
gradient = np.linspace(224, 32, image_width)
|
||||
image = np.tile(gradient, (image_height, 1))
|
||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
||||
return np2tensor(input_bg)
|
||||
case 'Right Light':
|
||||
gradient = np.linspace(32, 224, image_width)
|
||||
image = np.tile(gradient, (image_height, 1))
|
||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
||||
return np2tensor(input_bg)
|
||||
case 'Top Light':
|
||||
gradient = np.linspace(224, 32, image_height)[:, None]
|
||||
image = np.tile(gradient, (1, image_width))
|
||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
||||
return np2tensor(input_bg)
|
||||
case 'Bottom Light':
|
||||
gradient = np.linspace(32, 224, image_height)[:, None]
|
||||
image = np.tile(gradient, (1, image_width))
|
||||
input_bg = np.stack((image,) * 3, axis=-1).astype(np.uint8)
|
||||
return np2tensor(input_bg)
|
||||
case _:
|
||||
image = pil2tensor(Image.new('RGB', (1, 1), (0, 0, 0)))
|
||||
return image
|
||||
|
||||
|
||||
def apply(self, ic_model_path, model: ModelPatcher, c_concat: dict, ic_model=None) -> Tuple[ModelPatcher]:
|
||||
try:
|
||||
ModelPatcher.calculate_weight = calculate_weight_adjust_channel(ModelPatcher.calculate_weight)
|
||||
except:
|
||||
pass
|
||||
|
||||
device = comfy.model_management.get_torch_device()
|
||||
dtype = comfy.model_management.unet_dtype()
|
||||
work_model = model.clone()
|
||||
|
||||
# Apply scale factor.
|
||||
base_model: BaseModel = work_model.model
|
||||
scale_factor = base_model.model_config.latent_format.scale_factor
|
||||
|
||||
# [B, 4, H, W]
|
||||
concat_conds: torch.Tensor = c_concat["samples"] * scale_factor
|
||||
# [1, 4 * B, H, W]
|
||||
concat_conds = torch.cat([c[None, ...] for c in concat_conds], dim=1)
|
||||
|
||||
def unet_dummy_apply(unet_apply: Callable, params: UnetParams):
|
||||
"""A dummy unet apply wrapper serving as the endpoint of wrapper
|
||||
chain."""
|
||||
return unet_apply(x=params["input"], t=params["timestep"], **params["c"])
|
||||
|
||||
existing_wrapper = work_model.model_options.get(
|
||||
"model_function_wrapper", unet_dummy_apply
|
||||
)
|
||||
|
||||
def wrapper_func(unet_apply: Callable, params: UnetParams):
|
||||
return existing_wrapper(unet_apply, params=self.apply_c_concat(params, concat_conds))
|
||||
|
||||
work_model.set_model_unet_function_wrapper(wrapper_func)
|
||||
if not ic_model:
|
||||
ic_model = load_unet(ic_model_path)
|
||||
ic_model_state_dict = ic_model.model.diffusion_model.state_dict()
|
||||
|
||||
work_model.add_patches(
|
||||
patches={
|
||||
("diffusion_model." + key): (value.to(dtype=dtype, device=device),)
|
||||
for key, value in ic_model_state_dict.items()
|
||||
}
|
||||
)
|
||||
|
||||
return (work_model, ic_model)
|
||||
+936
-33
File diff suppressed because it is too large
Load Diff
@@ -53,20 +53,23 @@ class LayerDiffuse:
|
||||
sd_version = get_sd_version(model)
|
||||
model_url = LAYER_DIFFUSION[method.value][sd_version]["model_url"]
|
||||
|
||||
if image is not None:
|
||||
image = image.movedim(-1, 1)
|
||||
|
||||
try:
|
||||
ModelPatcher.calculate_weight = calculate_weight_adjust_channel(ModelPatcher.calculate_weight)
|
||||
except:
|
||||
pass
|
||||
|
||||
if method in [LayerMethod.FG_ONLY_CONV, LayerMethod.FG_ONLY_ATTN] and sd_version == 'sd15':
|
||||
if method in [LayerMethod.FG_ONLY_CONV, LayerMethod.FG_ONLY_ATTN] and sd_version == 'sd1':
|
||||
self.frames = 1
|
||||
elif method in [LayerMethod.BG_TO_BLEND, LayerMethod.FG_TO_BLEND, LayerMethod.BG_BLEND_TO_FG, LayerMethod.FG_BLEND_TO_BG] and sd_version == 'sd15':
|
||||
elif method in [LayerMethod.BG_TO_BLEND, LayerMethod.FG_TO_BLEND, LayerMethod.BG_BLEND_TO_FG, LayerMethod.FG_BLEND_TO_BG] and sd_version == 'sd1':
|
||||
self.frames = 2
|
||||
batch_size, _, height, width = samples['samples'].shape
|
||||
if batch_size % 2 != 0:
|
||||
raise Exception(f"The batch size should be a multiple of 2. 批次大小需为2的倍数")
|
||||
control_img = image
|
||||
elif method == LayerMethod.EVERYTHING and sd_version == 'sd15':
|
||||
elif method == LayerMethod.EVERYTHING and sd_version == 'sd1':
|
||||
batch_size, _, height, width = samples['samples'].shape
|
||||
self.frames = 3
|
||||
if batch_size % 3 != 0:
|
||||
@@ -77,7 +80,7 @@ class LayerDiffuse:
|
||||
model_path = get_local_filepath(model_url, LAYER_DIFFUSION_DIR)
|
||||
layer_lora_state_dict = load_layer_model_state_dict(model_path)
|
||||
work_model = model.clone()
|
||||
if sd_version == 'sd15':
|
||||
if sd_version == 'sd1':
|
||||
patcher = AttentionSharingPatcher(
|
||||
work_model, self.frames, use_control=control_img is not None
|
||||
)
|
||||
@@ -97,7 +100,7 @@ class LayerDiffuse:
|
||||
else:
|
||||
c_concat = model.model.latent_format.process_in(torch.cat([samples["samples"], blend_samples["samples"]], dim=1))
|
||||
samp_model, positive, negative = (work_model,) + self.apply_layer_c_concat(positive, negative, c_concat)
|
||||
elif sd_version == 'sd15':
|
||||
elif sd_version == 'sd1':
|
||||
if method in [LayerMethod.BG_TO_BLEND, LayerMethod.BG_BLEND_TO_FG]:
|
||||
additional_cond = (additional_cond[0], None)
|
||||
elif method in [LayerMethod.FG_TO_BLEND, LayerMethod.FG_BLEND_TO_BG]:
|
||||
@@ -166,10 +169,10 @@ class LayerDiffuse:
|
||||
alpha = []
|
||||
if layer_diffusion_method is not None:
|
||||
sd_version = get_sd_version(model)
|
||||
if sd_version not in ['sdxl', 'sd15']:
|
||||
if sd_version not in ['sdxl', 'sd1']:
|
||||
raise Exception(f"Only SDXL and SD1.5 model supported for Layer Diffusion")
|
||||
method = self.get_layer_diffusion_method(layer_diffusion_method, blend_samples is not None)
|
||||
sd15_allow = True if sd_version == 'sd15' and method in [LayerMethod.FG_ONLY_ATTN, LayerMethod.EVERYTHING, LayerMethod.BG_TO_BLEND, LayerMethod.BG_BLEND_TO_FG] else False
|
||||
sd15_allow = True if sd_version == 'sd1' and method in [LayerMethod.FG_ONLY_ATTN, LayerMethod.EVERYTHING, LayerMethod.BG_TO_BLEND, LayerMethod.BG_BLEND_TO_FG] else False
|
||||
sdxl_allow = True if sd_version == 'sdxl' and method in [LayerMethod.FG_ONLY_CONV, LayerMethod.FG_ONLY_ATTN, LayerMethod.BG_BLEND_TO_FG] else False
|
||||
if sdxl_allow or sd15_allow:
|
||||
if self.vae_transparent_decoder is None:
|
||||
|
||||
@@ -7,18 +7,17 @@ import comfy.model_management
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from tqdm import tqdm
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from ..libs.utils import install_package
|
||||
from packaging import version
|
||||
|
||||
try:
|
||||
install_package("diffusers", "0.27.2", True, "0.25.0")
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers import __version__
|
||||
if __version__:
|
||||
try:
|
||||
diffusers_version = float(__version__.replace('.', '').replace('dev','.'))
|
||||
except ValueError:
|
||||
diffusers_version = 270
|
||||
if diffusers_version < 270:
|
||||
if version.parse(__version__) < version.parse("0.26.0"):
|
||||
from diffusers.models.unet_2d_blocks import UNetMidBlock2D, get_down_block, get_up_block
|
||||
else:
|
||||
from diffusers.models.unets.unet_2d_blocks import UNetMidBlock2D, get_down_block, get_up_block
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
import urllib.parse
|
||||
from os import PathLike
|
||||
from aiohttp import web
|
||||
from aiohttp.web_urldispatcher import AbstractRoute, UrlDispatcher
|
||||
from server import PromptServer
|
||||
from pathlib import Path
|
||||
|
||||
# 文件限制大小(MB)
|
||||
max_size = 50
|
||||
def suffix_limiter(self: web.StaticResource, request: web.Request):
|
||||
suffixes = {".jpg", ".jpeg", ".png", ".gif", ".webp", ".bmp", ".tiff", ".svg", ".ico", ".apng", ".tif", ".hdr", ".exr"}
|
||||
rel_url = request.match_info["filename"]
|
||||
try:
|
||||
filename = Path(rel_url)
|
||||
if filename.anchor:
|
||||
raise web.HTTPForbidden()
|
||||
filepath = self._directory.joinpath(filename).resolve()
|
||||
if filepath.exists() and filepath.suffix.lower() not in suffixes:
|
||||
raise web.HTTPForbidden(reason="File type is not allowed")
|
||||
finally:
|
||||
pass
|
||||
|
||||
def filesize_limiter(self: web.StaticResource, request: web.Request):
|
||||
rel_url = request.match_info["filename"]
|
||||
try:
|
||||
filename = Path(rel_url)
|
||||
filepath = self._directory.joinpath(filename).resolve()
|
||||
if filepath.exists() and filepath.stat().st_size > max_size * 1024 * 1024:
|
||||
raise web.HTTPForbidden(reason="File size is too large")
|
||||
finally:
|
||||
pass
|
||||
class LimitResource(web.StaticResource):
|
||||
limiters = []
|
||||
|
||||
def push_limiter(self, limiter):
|
||||
self.limiters.append(limiter)
|
||||
|
||||
async def _handle(self, request: web.Request) -> web.StreamResponse:
|
||||
try:
|
||||
for limiter in self.limiters:
|
||||
limiter(self, request)
|
||||
except (ValueError, FileNotFoundError) as error:
|
||||
raise web.HTTPNotFound() from error
|
||||
|
||||
return await super()._handle(request)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
name = "'" + self.name + "'" if self.name is not None else ""
|
||||
return f'<LimitResource {name} {self._prefix} -> {self._directory!r}>'
|
||||
|
||||
class LimitRouter(web.StaticDef):
|
||||
def __repr__(self) -> str:
|
||||
info = []
|
||||
for name, value in sorted(self.kwargs.items()):
|
||||
info.append(f", {name}={value!r}")
|
||||
return f'<LimitRouter {self.prefix} -> {self.path}{"".join(info)}>'
|
||||
|
||||
def register(self, router: UrlDispatcher) -> list[AbstractRoute]:
|
||||
# resource = router.add_static(self.prefix, self.path, **self.kwargs)
|
||||
def add_static(
|
||||
self: UrlDispatcher,
|
||||
prefix: str,
|
||||
path: PathLike,
|
||||
*,
|
||||
name=None,
|
||||
expect_handler=None,
|
||||
chunk_size: int = 256 * 1024,
|
||||
show_index: bool = False,
|
||||
follow_symlinks: bool = False,
|
||||
append_version: bool = False,
|
||||
) -> web.AbstractResource:
|
||||
assert prefix.startswith("/")
|
||||
if prefix.endswith("/"):
|
||||
prefix = prefix[:-1]
|
||||
resource = LimitResource(
|
||||
prefix,
|
||||
path,
|
||||
name=name,
|
||||
expect_handler=expect_handler,
|
||||
chunk_size=chunk_size,
|
||||
show_index=show_index,
|
||||
follow_symlinks=follow_symlinks,
|
||||
append_version=append_version,
|
||||
)
|
||||
resource.push_limiter(suffix_limiter)
|
||||
resource.push_limiter(filesize_limiter)
|
||||
self.register_resource(resource)
|
||||
return resource
|
||||
resource = add_static(router, self.prefix, self.path, **self.kwargs)
|
||||
routes = resource.get_info().get("routes", {})
|
||||
return list(routes.values())
|
||||
|
||||
def path_to_url(path):
|
||||
if not path:
|
||||
return path
|
||||
path = path.replace("\\", "/")
|
||||
if not path.startswith("/"):
|
||||
path = "/" + path
|
||||
while path.startswith("//"):
|
||||
path = path[1:]
|
||||
path = path.replace("//", "/")
|
||||
return path
|
||||
|
||||
def add_static_resource(prefix, path,limit=False):
|
||||
app = PromptServer.instance.app
|
||||
prefix = path_to_url(prefix)
|
||||
prefix = urllib.parse.quote(prefix)
|
||||
prefix = path_to_url(prefix)
|
||||
if limit:
|
||||
route = LimitRouter(prefix, path, {"follow_symlinks": True})
|
||||
else:
|
||||
route = web.static(prefix, path, follow_symlinks=True)
|
||||
app.add_routes([route])
|
||||
@@ -4,9 +4,7 @@ import itertools
|
||||
|
||||
from comfy import model_management
|
||||
from comfy.sdxl_clip import SDXLClipModel, SDXLRefinerClipModel, SDXLClipG
|
||||
from nodes import NODE_CLASS_MAPPINGS, ConditioningConcat, CLIPTextEncode
|
||||
|
||||
from .libs.utils import compare_revision
|
||||
from nodes import NODE_CLASS_MAPPINGS, ConditioningConcat
|
||||
|
||||
def _grouper(n, iterable):
|
||||
it = iter(iterable)
|
||||
@@ -237,17 +235,17 @@ def encode_token_weights_g(model, token_weight_pairs):
|
||||
|
||||
|
||||
def encode_token_weights_l(model, token_weight_pairs):
|
||||
l_out, _ = model.clip_l.encode_token_weights(token_weight_pairs)
|
||||
return l_out, None
|
||||
l_out, pooled = model.clip_l.encode_token_weights(token_weight_pairs)
|
||||
return l_out, pooled
|
||||
|
||||
|
||||
def encode_token_weights(model, token_weight_pairs, encode_func):
|
||||
if model.layer_idx is not None:
|
||||
# 2016 [c2cb8e88] 及以上版本去除了sdxl clip的clip_layer方法
|
||||
if compare_revision(2016):
|
||||
model.cond_stage_model.set_clip_options({'layer': model.layer_idx})
|
||||
else:
|
||||
model.cond_stage_model.clip_layer(model.layer_idx)
|
||||
# if compare_revision(2016):
|
||||
model.cond_stage_model.set_clip_options({'layer': model.layer_idx})
|
||||
# else:
|
||||
# model.cond_stage_model.clip_layer(model.layer_idx)
|
||||
|
||||
model_management.load_model_gpu(model.patcher)
|
||||
return encode_func(model.cond_stage_model, token_weight_pairs)
|
||||
@@ -308,15 +306,16 @@ def advanced_encode(clip, text, token_normalization, weight_interpretation, w_ma
|
||||
|
||||
embeddings_final, pooled = prepareXL(embs_l, embs_g, pooled, clip_balance)
|
||||
|
||||
cond = [[embeddings_final,
|
||||
{"pooled_output": pooled, "width": width, "height": height, "crop_w": crop_w,
|
||||
"crop_h": crop_h, "target_width": target_width, "target_height": target_height}]]
|
||||
cond = [[embeddings_final, {"pooled_output": pooled}]]
|
||||
# cond = [[embeddings_final,
|
||||
# {"pooled_output": pooled, "width": width, "height": height, "crop_w": crop_w,
|
||||
# "crop_h": crop_h, "target_width": target_width, "target_height": target_height}]]
|
||||
else:
|
||||
embeddings_final, pooled = advanced_encode_from_tokens(tokenized['l'],
|
||||
token_normalization,
|
||||
weight_interpretation,
|
||||
lambda x: (clip.encode_from_tokens({'l': x}), None),
|
||||
w_max=w_max)
|
||||
lambda x: encode_token_weights(clip, x, encode_token_weights_l),
|
||||
w_max=w_max,return_pooled=True,)
|
||||
cond = [[embeddings_final, {"pooled_output": pooled}]]
|
||||
|
||||
if conditioning is not None:
|
||||
+79
-5
@@ -1,12 +1,86 @@
|
||||
cache = {}
|
||||
import itertools
|
||||
from typing import Optional
|
||||
|
||||
class TaggedCache:
|
||||
def __init__(self, tag_settings: Optional[dict]=None):
|
||||
self._tag_settings = tag_settings or {} # tag cache size
|
||||
self._data = {}
|
||||
|
||||
def __getitem__(self, key):
|
||||
for tag_data in self._data.values():
|
||||
if key in tag_data:
|
||||
return tag_data[key]
|
||||
raise KeyError(f'Key `{key}` does not exist')
|
||||
|
||||
def __setitem__(self, key, value: tuple):
|
||||
# value: (tag: str, (islist: bool, data: *))
|
||||
|
||||
# if key already exists, pop old value
|
||||
for tag_data in self._data.values():
|
||||
if key in tag_data:
|
||||
tag_data.pop(key, None)
|
||||
break
|
||||
|
||||
tag = value[0]
|
||||
if tag not in self._data:
|
||||
|
||||
try:
|
||||
from cachetools import LRUCache
|
||||
|
||||
default_size = 20
|
||||
if 'ckpt' in tag:
|
||||
default_size = 5
|
||||
elif tag in ['latent', 'image']:
|
||||
default_size = 100
|
||||
|
||||
self._data[tag] = LRUCache(maxsize=self._tag_settings.get(tag, default_size))
|
||||
|
||||
except (ImportError, ModuleNotFoundError):
|
||||
# TODO: implement a simple lru dict
|
||||
self._data[tag] = {}
|
||||
self._data[tag][key] = value
|
||||
|
||||
def __delitem__(self, key):
|
||||
for tag_data in self._data.values():
|
||||
if key in tag_data:
|
||||
del tag_data[key]
|
||||
return
|
||||
raise KeyError(f'Key `{key}` does not exist')
|
||||
|
||||
def __contains__(self, key):
|
||||
return any(key in tag_data for tag_data in self._data.values())
|
||||
|
||||
def items(self):
|
||||
yield from itertools.chain(*map(lambda x :x.items(), self._data.values()))
|
||||
|
||||
def get(self, key, default=None):
|
||||
"""D.get(k[,d]) -> D[k] if k in D, else d. d defaults to None."""
|
||||
for tag_data in self._data.values():
|
||||
if key in tag_data:
|
||||
return tag_data[key]
|
||||
return default
|
||||
|
||||
def clear(self):
|
||||
# clear all cache
|
||||
self._data = {}
|
||||
|
||||
cache_settings = {}
|
||||
cache = TaggedCache(cache_settings)
|
||||
cache_count = {}
|
||||
|
||||
|
||||
def update_cache(k, v):
|
||||
cache[k] = v
|
||||
def update_cache(k, tag, v):
|
||||
cache[k] = (tag, v)
|
||||
cnt = cache_count.get(k)
|
||||
if cnt is None:
|
||||
cnt = 0
|
||||
cache_count[k] = cnt
|
||||
else:
|
||||
cache_count[k] += 1
|
||||
cache_count[k] += 1
|
||||
def remove_cache(key):
|
||||
global cache
|
||||
if key == '*':
|
||||
cache = TaggedCache(cache_settings)
|
||||
elif key in cache:
|
||||
del cache[key]
|
||||
else:
|
||||
print(f"invalid {key}")
|
||||
@@ -0,0 +1,52 @@
|
||||
from server import PromptServer
|
||||
from aiohttp import web
|
||||
import time
|
||||
|
||||
class ChooserCancelled(Exception):
|
||||
pass
|
||||
|
||||
class ChooserMessage:
|
||||
stash = {}
|
||||
messages = {}
|
||||
cancelled = False
|
||||
|
||||
@classmethod
|
||||
def addMessage(cls, id, message):
|
||||
if message == '__cancel__':
|
||||
cls.messages = {}
|
||||
cls.cancelled = True
|
||||
elif message == '__start__':
|
||||
cls.messages = {}
|
||||
cls.stash = {}
|
||||
cls.cancelled = False
|
||||
else:
|
||||
cls.messages[str(id)] = message
|
||||
|
||||
@classmethod
|
||||
def waitForMessage(cls, id, period=0.1, asList=False):
|
||||
sid = str(id)
|
||||
while not (sid in cls.messages) and not ("-1" in cls.messages):
|
||||
if cls.cancelled:
|
||||
cls.cancelled = False
|
||||
raise ChooserCancelled()
|
||||
time.sleep(period)
|
||||
if cls.cancelled:
|
||||
cls.cancelled = False
|
||||
raise ChooserCancelled()
|
||||
message = cls.messages.pop(str(id), None) or cls.messages.pop("-1")
|
||||
try:
|
||||
if asList:
|
||||
return [int(x.strip()) for x in message.split(",")]
|
||||
else:
|
||||
return int(message.strip())
|
||||
except ValueError:
|
||||
print(
|
||||
f"ERROR IN IMAGE_CHOOSER - failed to parse '${message}' as ${'comma separated list of ints' if asList else 'int'}")
|
||||
return [1] if asList else 1
|
||||
|
||||
|
||||
@PromptServer.instance.routes.post('/easyuse/image_chooser_message')
|
||||
async def make_image_selection(request):
|
||||
post = await request.post()
|
||||
ChooserMessage.addMessage(post.get("id"), post.get("message"))
|
||||
return web.json_response({})
|
||||
@@ -0,0 +1,115 @@
|
||||
import torch
|
||||
from PIL import Image
|
||||
from torch import Tensor
|
||||
from torch.nn import functional as F
|
||||
|
||||
from torchvision.transforms import ToTensor, ToPILImage
|
||||
|
||||
def adain_color_fix(target: Image, source: Image):
|
||||
# Convert images to tensors
|
||||
to_tensor = ToTensor()
|
||||
target_tensor = to_tensor(target).unsqueeze(0)
|
||||
source_tensor = to_tensor(source).unsqueeze(0)
|
||||
|
||||
# Apply adaptive instance normalization
|
||||
result_tensor = adaptive_instance_normalization(target_tensor, source_tensor)
|
||||
|
||||
# Convert tensor back to image
|
||||
to_image = ToPILImage()
|
||||
result_image = to_image(result_tensor.squeeze(0).clamp_(0.0, 1.0))
|
||||
|
||||
return result_image
|
||||
|
||||
def wavelet_color_fix(target: Image, source: Image):
|
||||
source = source.resize(target.size, resample=Image.Resampling.LANCZOS)
|
||||
|
||||
# Convert images to tensors
|
||||
to_tensor = ToTensor()
|
||||
target_tensor = to_tensor(target).unsqueeze(0)
|
||||
source_tensor = to_tensor(source).unsqueeze(0)
|
||||
|
||||
# Apply wavelet reconstruction
|
||||
result_tensor = wavelet_reconstruction(target_tensor, source_tensor)
|
||||
|
||||
# Convert tensor back to image
|
||||
to_image = ToPILImage()
|
||||
result_image = to_image(result_tensor.squeeze(0).clamp_(0.0, 1.0))
|
||||
|
||||
return result_image
|
||||
|
||||
def calc_mean_std(feat: Tensor, eps=1e-5):
|
||||
"""Calculate mean and std for adaptive_instance_normalization.
|
||||
Args:
|
||||
feat (Tensor): 4D tensor.
|
||||
eps (float): A small value added to the variance to avoid
|
||||
divide-by-zero. Default: 1e-5.
|
||||
"""
|
||||
size = feat.size()
|
||||
assert len(size) == 4, 'The input feature should be 4D tensor.'
|
||||
b, c = size[:2]
|
||||
feat_var = feat.view(b, c, -1).var(dim=2) + eps
|
||||
feat_std = feat_var.sqrt().view(b, c, 1, 1)
|
||||
feat_mean = feat.view(b, c, -1).mean(dim=2).view(b, c, 1, 1)
|
||||
return feat_mean, feat_std
|
||||
|
||||
def adaptive_instance_normalization(content_feat:Tensor, style_feat:Tensor):
|
||||
"""Adaptive instance normalization.
|
||||
Adjust the reference features to have the similar color and illuminations
|
||||
as those in the degradate features.
|
||||
Args:
|
||||
content_feat (Tensor): The reference feature.
|
||||
style_feat (Tensor): The degradate features.
|
||||
"""
|
||||
size = content_feat.size()
|
||||
style_mean, style_std = calc_mean_std(style_feat)
|
||||
content_mean, content_std = calc_mean_std(content_feat)
|
||||
normalized_feat = (content_feat - content_mean.expand(size)) / content_std.expand(size)
|
||||
return normalized_feat * style_std.expand(size) + style_mean.expand(size)
|
||||
|
||||
def wavelet_blur(image: Tensor, radius: int):
|
||||
"""
|
||||
Apply wavelet blur to the input tensor.
|
||||
"""
|
||||
# input shape: (1, 3, H, W)
|
||||
# convolution kernel
|
||||
kernel_vals = [
|
||||
[0.0625, 0.125, 0.0625],
|
||||
[0.125, 0.25, 0.125],
|
||||
[0.0625, 0.125, 0.0625],
|
||||
]
|
||||
kernel = torch.tensor(kernel_vals, dtype=image.dtype, device=image.device)
|
||||
# add channel dimensions to the kernel to make it a 4D tensor
|
||||
kernel = kernel[None, None]
|
||||
# repeat the kernel across all input channels
|
||||
kernel = kernel.repeat(3, 1, 1, 1)
|
||||
image = F.pad(image, (radius, radius, radius, radius), mode='replicate')
|
||||
# apply convolution
|
||||
output = F.conv2d(image, kernel, groups=3, dilation=radius)
|
||||
return output
|
||||
|
||||
def wavelet_decomposition(image: Tensor, levels=5):
|
||||
"""
|
||||
Apply wavelet decomposition to the input tensor.
|
||||
This function only returns the low frequency & the high frequency.
|
||||
"""
|
||||
high_freq = torch.zeros_like(image)
|
||||
for i in range(levels):
|
||||
radius = 2 ** i
|
||||
low_freq = wavelet_blur(image, radius)
|
||||
high_freq += (image - low_freq)
|
||||
image = low_freq
|
||||
|
||||
return high_freq, low_freq
|
||||
|
||||
def wavelet_reconstruction(content_feat:Tensor, style_feat:Tensor):
|
||||
"""
|
||||
Apply wavelet decomposition, so that the content will have the same color as the style.
|
||||
"""
|
||||
# calculate the wavelet decomposition of the content feature
|
||||
content_high_freq, content_low_freq = wavelet_decomposition(content_feat)
|
||||
del content_low_freq
|
||||
# calculate the wavelet decomposition of the style feature
|
||||
style_high_freq, style_low_freq = wavelet_decomposition(style_feat)
|
||||
del style_high_freq
|
||||
# reconstruct the content feature with the style's high frequency
|
||||
return content_high_freq + style_low_freq
|
||||
+13
-7
@@ -1,14 +1,20 @@
|
||||
from .utils import find_wildcards_seed, find_nearest_steps, is_linked_styles_selector
|
||||
from ..log import log_node_warn
|
||||
from ..adv_encode import advanced_encode
|
||||
from ..wildcards import process_with_loras
|
||||
from .log import log_node_warn
|
||||
from .translate import zh_to_en, has_chinese
|
||||
from .wildcards import process_with_loras
|
||||
from .adv_encode import advanced_encode
|
||||
|
||||
from nodes import ConditioningConcat, ConditioningCombine, ConditioningAverage, ConditioningSetTimestepRange
|
||||
|
||||
def prompt_to_cond(type, model, clip, clip_skip, lora_stack, text, prompt_token_normalization, prompt_weight_interpretation, a1111_prompt_style ,my_unique_id, prompt, easyCache, can_load_lora=True):
|
||||
def prompt_to_cond(type, model, clip, clip_skip, lora_stack, text, prompt_token_normalization, prompt_weight_interpretation, a1111_prompt_style ,my_unique_id, prompt, easyCache, can_load_lora=True, steps=None):
|
||||
styles_selector = is_linked_styles_selector(prompt, my_unique_id, type)
|
||||
title = "正面提示词" if type == 'positive' else "负面提示词"
|
||||
log_node_warn("正在处理" + title + "...")
|
||||
log_node_warn("正在进行" + title + "...")
|
||||
|
||||
# Translate cn to en
|
||||
if has_chinese(text):
|
||||
text = zh_to_en([text])[0]
|
||||
|
||||
positive_seed = find_wildcards_seed(my_unique_id, text, prompt)
|
||||
model, clip, text, cond_decode, show_prompt, pipe_lora_stack = process_with_loras(
|
||||
text, model, clip, type, positive_seed, can_load_lora, lora_stack, easyCache)
|
||||
@@ -18,8 +24,8 @@ def prompt_to_cond(type, model, clip, clip_skip, lora_stack, text, prompt_token_
|
||||
if clip_skip != 0:
|
||||
clipped.clip_layer(clip_skip)
|
||||
|
||||
log_node_warn("正在处理" + title + "编码...")
|
||||
steps = find_nearest_steps(my_unique_id, prompt)
|
||||
log_node_warn("正在进行" + title + "编码...")
|
||||
steps = steps if steps is not None else find_nearest_steps(my_unique_id, prompt)
|
||||
return (advanced_encode(clipped, text, prompt_token_normalization,
|
||||
prompt_weight_interpretation, w_max=1.0,
|
||||
apply_to_pooled='enable',
|
||||
|
||||
+3
-17
@@ -7,26 +7,12 @@ class easyControlnet:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def load_controlnet(self, control_net_name, control_net, scale_soft_weights):
|
||||
if control_net is None:
|
||||
if scale_soft_weights < 1:
|
||||
if "ScaledSoftControlNetWeights" in NODE_CLASS_MAPPINGS:
|
||||
soft_weight_cls = NODE_CLASS_MAPPINGS['ScaledSoftControlNetWeights']
|
||||
(weights, timestep_keyframe) = soft_weight_cls().load_weights(scale_soft_weights, False)
|
||||
cn_adv_cls = NODE_CLASS_MAPPINGS['ControlNetLoaderAdvanced']
|
||||
control_net, = cn_adv_cls().load_controlnet(control_net_name, timestep_keyframe)
|
||||
else:
|
||||
raise Exception(f"[Advanced-ControlNet Not Found] you need to install 'COMFYUI-Advanced-ControlNet'")
|
||||
else:
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
||||
control_net = comfy.controlnet.load_controlnet(controlnet_path)
|
||||
return control_net
|
||||
|
||||
def apply(self, control_net_name, image, positive, negative, strength, start_percent=0, end_percent=1, control_net=None, scale_soft_weights=1, mask=None):
|
||||
def apply(self, control_net_name, image, positive, negative, strength, start_percent=0, end_percent=1, control_net=None, scale_soft_weights=1, mask=None, easyCache=None, use_cache=True):
|
||||
if strength == 0:
|
||||
return (positive, negative)
|
||||
|
||||
control_net = self.load_controlnet(control_net_name, control_net, scale_soft_weights)
|
||||
if control_net is None:
|
||||
control_net = easyCache.load_controlnet(control_net_name, scale_soft_weights, use_cache)
|
||||
|
||||
if mask is not None:
|
||||
mask = mask.to(self.device)
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
@staticmethod
|
||||
def easyIn(t: float)-> float:
|
||||
return t*t
|
||||
@staticmethod
|
||||
def easyOut(t: float)-> float:
|
||||
return -(t * (t - 2))
|
||||
@staticmethod
|
||||
def easyInOut(t: float)-> float:
|
||||
if t < 0.5:
|
||||
return 2*t*t
|
||||
else:
|
||||
return (-2*t*t) + (4*t) - 1
|
||||
|
||||
class EasingBase:
|
||||
|
||||
def easing(self, t: float, function='linear') -> float:
|
||||
if function == 'easyIn':
|
||||
return easyIn(t)
|
||||
elif function == 'easyOut':
|
||||
return easyOut(t)
|
||||
elif function == 'easyInOut':
|
||||
return easyInOut(t)
|
||||
else:
|
||||
return t
|
||||
|
||||
def ease(self, start, end, t) -> float:
|
||||
return end * t + start * (1 - t)
|
||||
@@ -19,6 +19,8 @@ class InpaintWorker:
|
||||
def __init__(self, node_name):
|
||||
self.node_name = node_name if node_name is not None else ""
|
||||
self.original_calculate_weight = ModelPatcher.calculate_weight
|
||||
if not hasattr(ModelPatcher, "original_calculate_weight"):
|
||||
ModelPatcher.original_calculate_weight = self.original_calculate_weight
|
||||
self.injected_model_patcher_calculate_weight = False
|
||||
|
||||
def load_fooocus_patch(self, lora: dict, to_load: dict):
|
||||
@@ -0,0 +1,194 @@
|
||||
import os
|
||||
import base64
|
||||
import torch
|
||||
import numpy as np
|
||||
from enum import Enum
|
||||
from PIL import Image
|
||||
from io import BytesIO
|
||||
from typing import List, Union
|
||||
|
||||
import folder_paths
|
||||
from .utils import install_package
|
||||
|
||||
# PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
# np to Tensor
|
||||
def np2tensor(img_np: Union[np.ndarray, List[np.ndarray]]) -> torch.Tensor:
|
||||
if isinstance(img_np, list):
|
||||
return torch.cat([np2tensor(img) for img in img_np], dim=0)
|
||||
return torch.from_numpy(img_np.astype(np.float32) / 255.0).unsqueeze(0)
|
||||
# Tensor to np
|
||||
def tensor2np(tensor: torch.Tensor) -> List[np.ndarray]:
|
||||
if len(tensor.shape) == 3: # Single image
|
||||
return np.clip(255.0 * tensor.cpu().numpy(), 0, 255).astype(np.uint8)
|
||||
else: # Batch of images
|
||||
return [np.clip(255.0 * t.cpu().numpy(), 0, 255).astype(np.uint8) for t in tensor]
|
||||
|
||||
def pil2byte(pil_image, format='PNG'):
|
||||
byte_arr = BytesIO()
|
||||
pil_image.save(byte_arr, format=format)
|
||||
byte_arr.seek(0)
|
||||
return byte_arr
|
||||
|
||||
def image2base64(image_base64):
|
||||
image_bytes = base64.b64decode(image_base64)
|
||||
image_data = Image.open(BytesIO(image_bytes))
|
||||
return image_data
|
||||
|
||||
# Get new bounds
|
||||
def get_new_bounds(width, height, left, right, top, bottom):
|
||||
"""Returns the new bounds for an image with inset crop data."""
|
||||
left = 0 + left
|
||||
right = width - right
|
||||
top = 0 + top
|
||||
bottom = height - bottom
|
||||
return (left, right, top, bottom)
|
||||
|
||||
def RGB2RGBA(image: Image, mask: Image) -> Image:
|
||||
(R, G, B) = image.convert('RGB').split()
|
||||
return Image.merge('RGBA', (R, G, B, mask.convert('L')))
|
||||
|
||||
def image2mask(image: Image) -> torch.Tensor:
|
||||
_image = image.convert('RGBA')
|
||||
alpha = _image.split()[0]
|
||||
bg = Image.new("L", _image.size)
|
||||
_image = Image.merge('RGBA', (bg, bg, bg, alpha))
|
||||
ret_mask = torch.tensor([pil2tensor(_image)[0, :, :, 3].tolist()])
|
||||
return ret_mask
|
||||
|
||||
def mask2image(mask: torch.Tensor) -> Image:
|
||||
masks = tensor2np(mask)
|
||||
for m in masks:
|
||||
_mask = Image.fromarray(m).convert("L")
|
||||
_image = Image.new("RGBA", _mask.size, color='white')
|
||||
_image = Image.composite(
|
||||
_image, Image.new("RGBA", _mask.size, color='black'), _mask)
|
||||
return _image
|
||||
|
||||
# 图像融合
|
||||
class blendImage:
|
||||
def g(self, x):
|
||||
return torch.where(x <= 0.25, ((16 * x - 12) * x + 4) * x, torch.sqrt(x))
|
||||
|
||||
def blend_mode(self, img1, img2, mode):
|
||||
if mode == "normal":
|
||||
return img2
|
||||
elif mode == "multiply":
|
||||
return img1 * img2
|
||||
elif mode == "screen":
|
||||
return 1 - (1 - img1) * (1 - img2)
|
||||
elif mode == "overlay":
|
||||
return torch.where(img1 <= 0.5, 2 * img1 * img2, 1 - 2 * (1 - img1) * (1 - img2))
|
||||
elif mode == "soft_light":
|
||||
return torch.where(img2 <= 0.5, img1 - (1 - 2 * img2) * img1 * (1 - img1),
|
||||
img1 + (2 * img2 - 1) * (self.g(img1) - img1))
|
||||
elif mode == "difference":
|
||||
return img1 - img2
|
||||
else:
|
||||
raise ValueError(f"Unsupported blend mode: {mode}")
|
||||
|
||||
def blend_images(self, image1: torch.Tensor, image2: torch.Tensor, blend_factor: float, blend_mode: str = 'normal'):
|
||||
image2 = image2.to(image1.device)
|
||||
if image1.shape != image2.shape:
|
||||
image2 = image2.permute(0, 3, 1, 2)
|
||||
image2 = comfy.utils.common_upscale(image2, image1.shape[2], image1.shape[1], upscale_method='bicubic',
|
||||
crop='center')
|
||||
image2 = image2.permute(0, 2, 3, 1)
|
||||
|
||||
blended_image = self.blend_mode(image1, image2, blend_mode)
|
||||
blended_image = image1 * (1 - blend_factor) + blended_image * blend_factor
|
||||
blended_image = torch.clamp(blended_image, 0, 1)
|
||||
return blended_image
|
||||
|
||||
|
||||
|
||||
|
||||
class ResizeMode(Enum):
|
||||
RESIZE = "Just Resize"
|
||||
INNER_FIT = "Crop and Resize"
|
||||
OUTER_FIT = "Resize and Fill"
|
||||
def int_value(self):
|
||||
if self == ResizeMode.RESIZE:
|
||||
return 0
|
||||
elif self == ResizeMode.INNER_FIT:
|
||||
return 1
|
||||
elif self == ResizeMode.OUTER_FIT:
|
||||
return 2
|
||||
assert False, "NOTREACHED"
|
||||
|
||||
|
||||
|
||||
# CLIP反推
|
||||
import comfy.utils
|
||||
from torchvision import transforms
|
||||
Config, Interrogator = None, None
|
||||
class CI_Inference:
|
||||
ci_model = None
|
||||
cache_path: str
|
||||
|
||||
def __init__(self):
|
||||
self.ci_model = None
|
||||
self.low_vram = False
|
||||
self.cache_path = os.path.join(folder_paths.models_dir, "clip_interrogator")
|
||||
|
||||
def _load_model(self, model_name, low_vram=False):
|
||||
if not (self.ci_model and model_name == self.ci_model.config.clip_model_name and self.low_vram == low_vram):
|
||||
self.low_vram = low_vram
|
||||
print(f"Load model: {model_name}")
|
||||
|
||||
config = Config(
|
||||
device="cuda" if torch.cuda.is_available() else "cpu",
|
||||
download_cache=True,
|
||||
clip_model_name=model_name,
|
||||
clip_model_path=self.cache_path,
|
||||
cache_path=self.cache_path,
|
||||
caption_model_name='blip-large'
|
||||
)
|
||||
|
||||
if low_vram:
|
||||
config.apply_low_vram_defaults()
|
||||
|
||||
self.ci_model = Interrogator(config)
|
||||
|
||||
def _interrogate(self, image, mode, caption=None):
|
||||
if mode == 'best':
|
||||
prompt = self.ci_model.interrogate(image, caption=caption)
|
||||
elif mode == 'classic':
|
||||
prompt = self.ci_model.interrogate_classic(image, caption=caption)
|
||||
elif mode == 'fast':
|
||||
prompt = self.ci_model.interrogate_fast(image, caption=caption)
|
||||
elif mode == 'negative':
|
||||
prompt = self.ci_model.interrogate_negative(image)
|
||||
else:
|
||||
raise Exception(f"Unknown mode {mode}")
|
||||
return prompt
|
||||
|
||||
def image_to_prompt(self, image, mode, model_name='ViT-L-14/openai', low_vram=False):
|
||||
try:
|
||||
from clip_interrogator import Config, Interrogator
|
||||
global Config, Interrogator
|
||||
except:
|
||||
install_package("clip_interrogator", "0.6.0")
|
||||
from clip_interrogator import Config, Interrogator
|
||||
|
||||
pbar = comfy.utils.ProgressBar(len(image))
|
||||
|
||||
self._load_model(model_name, low_vram)
|
||||
prompt = []
|
||||
for i in range(len(image)):
|
||||
im = image[i]
|
||||
|
||||
im = tensor2pil(im)
|
||||
im = im.convert('RGB')
|
||||
|
||||
_prompt = self._interrogate(im, mode)
|
||||
pbar.update(1)
|
||||
prompt.append(_prompt)
|
||||
|
||||
return prompt
|
||||
|
||||
ci = CI_Inference()
|
||||
+142
-14
@@ -1,16 +1,21 @@
|
||||
import time, os, psutil
|
||||
import folder_paths
|
||||
import comfy.utils
|
||||
import comfy.sd
|
||||
import folder_paths
|
||||
import comfy.controlnet
|
||||
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from nodes import NODE_CLASS_MAPPINGS
|
||||
from collections import defaultdict
|
||||
from ..log import log_node_info, log_node_error
|
||||
from .log import log_node_info, log_node_error
|
||||
|
||||
stable_diffusion_loaders = ["easy a1111Loader", "easy comfyLoader", "easy zero123Loader", "easy svdLoader"]
|
||||
stable_diffusion_loaders = ["easy fullLoader", "easy a1111Loader", "easy comfyLoader", "easy zero123Loader", "easy svdLoader"]
|
||||
stable_cascade_loaders = ["easy cascadeLoader"]
|
||||
controlnet_loaders = ["easy controlnetLoader", "easy controlnetLoaderADV"]
|
||||
instant_loaders = ["easy instantIDApply", "easy instantIDApplyADV"]
|
||||
cascade_vae_node = ["easy preSamplingCascade", "easy fullCascadeKSampler"]
|
||||
model_merge_node = ["easy XYInputs: ModelMergeBlocks"]
|
||||
lora_widget = ["easy a1111Loader", "easy comfyLoader"]
|
||||
lora_widget = ["easy fullLoader", "easy a1111Loader", "easy comfyLoader"]
|
||||
|
||||
class easyLoader:
|
||||
def __init__(self):
|
||||
@@ -22,6 +27,7 @@ class easyLoader:
|
||||
"bvae": defaultdict(tuple),
|
||||
"vae": defaultdict(object),
|
||||
"lora": defaultdict(dict), # {lora_name: {UID: (model_lora, clip_lora)}}
|
||||
"controlnet": defaultdict(dict),
|
||||
}
|
||||
self.memory_threshold = self.determine_memory_threshold(0.7)
|
||||
self.lora_name_cache = []
|
||||
@@ -50,9 +56,17 @@ class easyLoader:
|
||||
for key in keys - desired_names:
|
||||
del self.loaded_objects[object_type][key]
|
||||
|
||||
def get_input_value(self, entry, key):
|
||||
def get_input_value(self, entry, key, prompt=None):
|
||||
val = entry["inputs"][key]
|
||||
return val if isinstance(val, str) else val[0]
|
||||
if isinstance(val, str):
|
||||
return val
|
||||
elif isinstance(val, list):
|
||||
if prompt is not None and val[0]:
|
||||
return prompt[val[0]]['inputs'][key]
|
||||
else:
|
||||
return val[0]
|
||||
else:
|
||||
return str(val)
|
||||
|
||||
def process_pipe_loader(self, entry, desired_ckpt_names, desired_vae_names, desired_lora_names, desired_lora_settings, num_loras=3, suffix=""):
|
||||
for idx in range(1, num_loras + 1):
|
||||
@@ -71,10 +85,10 @@ class easyLoader:
|
||||
desired_vae_names = set()
|
||||
desired_lora_names = set()
|
||||
desired_lora_settings = set()
|
||||
desired_controlnet_names = set()
|
||||
|
||||
for entry in prompt.values():
|
||||
class_type = entry["class_type"]
|
||||
|
||||
if class_type in lora_widget:
|
||||
lora_name = self.get_input_value(entry, "lora_name")
|
||||
desired_lora_names.add(lora_name)
|
||||
@@ -82,7 +96,7 @@ class easyLoader:
|
||||
desired_lora_settings.add(setting)
|
||||
|
||||
if class_type in stable_diffusion_loaders:
|
||||
desired_ckpt_names.add(self.get_input_value(entry, "ckpt_name"))
|
||||
desired_ckpt_names.add(self.get_input_value(entry, "ckpt_name", prompt))
|
||||
desired_vae_names.add(self.get_input_value(entry, "vae_name"))
|
||||
|
||||
elif class_type in stable_cascade_loaders:
|
||||
@@ -99,6 +113,16 @@ class easyLoader:
|
||||
if decode_vae_name and decode_vae_name != 'None':
|
||||
desired_vae_names.add(decode_vae_name)
|
||||
|
||||
elif class_type in controlnet_loaders:
|
||||
control_net_name = self.get_input_value(entry, "control_net_name", prompt)
|
||||
scale_soft_weights = self.get_input_value(entry, "scale_soft_weights")
|
||||
desired_controlnet_names.add(f'{control_net_name};{scale_soft_weights}')
|
||||
|
||||
elif class_type in instant_loaders:
|
||||
control_net_name = self.get_input_value(entry, "control_net_name", prompt)
|
||||
scale_soft_weights = self.get_input_value(entry, "cn_soft_weights")
|
||||
desired_controlnet_names.add(f'{control_net_name};{scale_soft_weights}')
|
||||
|
||||
elif class_type in model_merge_node:
|
||||
desired_ckpt_names.add(self.get_input_value(entry, "ckpt_name_1"))
|
||||
desired_ckpt_names.add(self.get_input_value(entry, "ckpt_name_2"))
|
||||
@@ -106,7 +130,7 @@ class easyLoader:
|
||||
if vae_use != 'Use Model 1' and vae_use != 'Use Model 2':
|
||||
desired_vae_names.add(vae_use)
|
||||
|
||||
object_types = ["ckpt", "unet", "clip", "bvae", "vae", "lora"]
|
||||
object_types = ["ckpt", "unet", "clip", "bvae", "vae", "lora", "controlnet"]
|
||||
for object_type in object_types:
|
||||
if object_type == 'unet':
|
||||
desired_names = desired_unet_names
|
||||
@@ -117,6 +141,8 @@ class easyLoader:
|
||||
desired_names = desired_ckpt_names
|
||||
elif object_type == "vae":
|
||||
desired_names = desired_vae_names
|
||||
elif object_type == "controlnet":
|
||||
desired_names = desired_controlnet_names
|
||||
else:
|
||||
desired_names = desired_lora_names
|
||||
self.clear_unused_objects(desired_names, object_type)
|
||||
@@ -155,7 +181,7 @@ class easyLoader:
|
||||
current_memory = self.get_memory_usage()
|
||||
if current_memory < self.memory_threshold:
|
||||
return
|
||||
eviction_order = ["vae", "lora", "bvae", "clip", "ckpt"]
|
||||
eviction_order = ["vae", "lora", "bvae", "clip", "ckpt", "controlnet"]
|
||||
for obj_type in eviction_order:
|
||||
if current_memory < self.memory_threshold:
|
||||
break
|
||||
@@ -225,7 +251,27 @@ class easyLoader:
|
||||
|
||||
return model
|
||||
|
||||
def load_clip(self, clip_name, type='stable_diffusion'):
|
||||
def load_controlnet(self, control_net_name, scale_soft_weights=1, use_cache=True):
|
||||
unique_id = f'{control_net_name};{str(scale_soft_weights)}'
|
||||
if use_cache and unique_id in self.loaded_objects["controlnet"]:
|
||||
return self.loaded_objects["controlnet"][unique_id][0]
|
||||
if scale_soft_weights < 1:
|
||||
if "ScaledSoftControlNetWeights" in NODE_CLASS_MAPPINGS:
|
||||
soft_weight_cls = NODE_CLASS_MAPPINGS['ScaledSoftControlNetWeights']
|
||||
(weights, timestep_keyframe) = soft_weight_cls().load_weights(scale_soft_weights, False)
|
||||
cn_adv_cls = NODE_CLASS_MAPPINGS['ControlNetLoaderAdvanced']
|
||||
control_net, = cn_adv_cls().load_controlnet(control_net_name, timestep_keyframe)
|
||||
else:
|
||||
raise Exception(
|
||||
f"[Advanced-ControlNet Not Found] you need to install 'COMFYUI-Advanced-ControlNet'")
|
||||
else:
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
||||
control_net = comfy.controlnet.load_controlnet(controlnet_path)
|
||||
if use_cache:
|
||||
self.add_to_cache("controlnet", unique_id, control_net)
|
||||
self.eviction_based_on_memory()
|
||||
return control_net
|
||||
def load_clip(self, clip_name, type='stable_diffusion', load_clip=None):
|
||||
if type == 'stable_diffusion':
|
||||
clip_type = comfy.sd.CLIPType.STABLE_DIFFUSION
|
||||
else:
|
||||
@@ -258,8 +304,6 @@ class easyLoader:
|
||||
orig_lora_name = lora_name
|
||||
lora_name = self.resolve_lora_name(lora_name)
|
||||
|
||||
|
||||
|
||||
if lora_name is not None:
|
||||
lora_path = folder_paths.get_full_path("loras", lora_name)
|
||||
else:
|
||||
@@ -278,6 +322,32 @@ class easyLoader:
|
||||
lbw_a, lbw_b, "", lbw)
|
||||
else:
|
||||
_lora = comfy.utils.load_torch_file(lora_path, safe_load=True)
|
||||
keys = _lora.keys()
|
||||
if "down_blocks.0.resnets.0.norm1.bias" in keys:
|
||||
print('Using LORA for Resadapter')
|
||||
key_map = {}
|
||||
key_map = comfy.lora.model_lora_keys_unet(model.model, key_map)
|
||||
mapping_norm = {}
|
||||
|
||||
for key in keys:
|
||||
if ".weight" in key:
|
||||
key_name_in_ori_sd = key_map[key.replace(".weight", "")]
|
||||
mapping_norm[key_name_in_ori_sd] = _lora[key]
|
||||
elif ".bias" in key:
|
||||
key_name_in_ori_sd = key_map[key.replace(".bias", "")]
|
||||
mapping_norm[key_name_in_ori_sd.replace(".weight", ".bias")] = _lora[
|
||||
key
|
||||
]
|
||||
else:
|
||||
print("===>Unexpected key", key)
|
||||
mapping_norm[key] = _lora[key]
|
||||
|
||||
for k in mapping_norm.keys():
|
||||
if k not in model.model.state_dict():
|
||||
print("===>Missing key:", k)
|
||||
model.model.load_state_dict(mapping_norm, strict=False)
|
||||
return (model, clip)
|
||||
|
||||
model, clip = comfy.sd.load_lora_for_models(model, clip, _lora, model_strength, clip_strength)
|
||||
|
||||
self.add_to_cache("lora", unique_id, (model, clip))
|
||||
@@ -306,4 +376,62 @@ class easyLoader:
|
||||
self.lora_name_cache.append(x)
|
||||
return x
|
||||
|
||||
return None
|
||||
return None
|
||||
|
||||
def load_main(self, ckpt_name, config_name, vae_name, lora_name, lora_model_strength, lora_clip_strength, optional_lora_stack, model_override, clip_override, vae_override, prompt):
|
||||
model: ModelPatcher | None = None
|
||||
clip: comfy.sd.CLIP | None = None
|
||||
vae: comfy.sd.VAE | None = None
|
||||
clip_vision = None
|
||||
lora_stack = []
|
||||
|
||||
can_load_lora = True
|
||||
# 判断是否存在 模型或Lora叠加xyplot, 若存在优先缓存第一个模型
|
||||
xy_model_id = next((x for x in prompt if str(prompt[x]["class_type"]) in ["easy XYInputs: ModelMergeBlocks",
|
||||
"easy XYInputs: Checkpoint"]), None)
|
||||
xy_lora_id = next((x for x in prompt if str(prompt[x]["class_type"]) == "easy XYInputs: Lora"), None)
|
||||
if xy_lora_id is not None:
|
||||
can_load_lora = False
|
||||
if xy_model_id is not None:
|
||||
node = prompt[xy_model_id]
|
||||
if "ckpt_name_1" in node["inputs"]:
|
||||
ckpt_name_1 = node["inputs"]["ckpt_name_1"]
|
||||
model, clip, vae, clip_vision = self.load_checkpoint(ckpt_name_1)
|
||||
can_load_lora = False
|
||||
# Load models
|
||||
elif model_override is not None and clip_override is not None and vae_override is not None:
|
||||
model = model_override
|
||||
clip = clip_override
|
||||
vae = vae_override
|
||||
elif model_override is not None:
|
||||
raise Exception(f"[ERROR] clip or vae is missing")
|
||||
elif vae_override is not None:
|
||||
raise Exception(f"[ERROR] model or clip is missing")
|
||||
elif clip_override is not None:
|
||||
raise Exception(f"[ERROR] model or vae is missing")
|
||||
else:
|
||||
model, clip, vae, clip_vision = self.load_checkpoint(ckpt_name, config_name)
|
||||
|
||||
if optional_lora_stack is not None and can_load_lora:
|
||||
for lora in optional_lora_stack:
|
||||
lora = {"lora_name": lora[0], "model": model, "clip": clip, "model_strength": lora[1],
|
||||
"clip_strength": lora[2]}
|
||||
model, clip = self.load_lora(lora)
|
||||
lora['model'] = model
|
||||
lora['clip'] = clip
|
||||
lora_stack.append(lora)
|
||||
|
||||
if lora_name != "None" and can_load_lora:
|
||||
lora = {"lora_name": lora_name, "model": model, "clip": clip, "model_strength": lora_model_strength,
|
||||
"clip_strength": lora_clip_strength}
|
||||
model, clip = self.load_lora(lora)
|
||||
lora_stack.append(lora)
|
||||
|
||||
# Check for custom VAE
|
||||
if vae_name not in ["Baked VAE", "Baked-VAE"]:
|
||||
vae = self.load_vae(vae_name)
|
||||
# CLIP skip
|
||||
if not clip:
|
||||
raise Exception("No CLIP found")
|
||||
|
||||
return model, clip, vae, clip_vision, lora_stack
|
||||
@@ -0,0 +1,58 @@
|
||||
import json
|
||||
import os
|
||||
import folder_paths
|
||||
import server
|
||||
from .utils import find_tags
|
||||
|
||||
class easyModelManager:
|
||||
|
||||
def __init__(self):
|
||||
self.img_suffixes = [".png", ".jpg", ".jpeg", ".gif", ".webp", ".bmp", ".tiff", ".svg", ".tif", ".tiff"]
|
||||
self.default_suffixes = [".ckpt", ".pt", ".bin", ".pth", ".safetensors"]
|
||||
self.models_config = {
|
||||
"checkpoints": {"suffix": self.default_suffixes},
|
||||
"loras": {"suffix": self.default_suffixes},
|
||||
"unet": {"suffix": self.default_suffixes},
|
||||
}
|
||||
self.model_lists = {}
|
||||
|
||||
def find_thumbnail(self, model_type, name):
|
||||
file_no_ext = os.path.splitext(name)[0]
|
||||
for ext in self.img_suffixes:
|
||||
full_path = folder_paths.get_full_path(model_type, file_no_ext + ext)
|
||||
if os.path.isfile(str(full_path)):
|
||||
return full_path
|
||||
return None
|
||||
|
||||
def get_model_lists(self, model_type):
|
||||
if model_type not in self.models_config:
|
||||
return []
|
||||
filenames = folder_paths.get_filename_list(model_type)
|
||||
model_lists = []
|
||||
for name in filenames:
|
||||
model_suffix = os.path.splitext(name)[-1]
|
||||
if model_suffix not in self.models_config[model_type]["suffix"]:
|
||||
continue
|
||||
else:
|
||||
cfg = {
|
||||
"name": os.path.basename(os.path.splitext(name)[0]),
|
||||
"full_name": name,
|
||||
"remark": '',
|
||||
"file_path": folder_paths.get_full_path(model_type, name),
|
||||
"type": model_type,
|
||||
"suffix": model_suffix,
|
||||
"dir_tags": find_tags(name),
|
||||
"cover": self.find_thumbnail(model_type, name),
|
||||
"metadata": None,
|
||||
"sha256": None
|
||||
}
|
||||
model_lists.append(cfg)
|
||||
|
||||
return model_lists
|
||||
|
||||
def get_model_info(self, model_type, model_name):
|
||||
pass
|
||||
|
||||
# if __name__ == "__main__":
|
||||
# manager = easyModelManager()
|
||||
# print(manager.get_model_lists("checkpoints"))
|
||||
+148
-8
@@ -5,13 +5,14 @@ import latent_preview
|
||||
from nodes import MAX_RESOLUTION
|
||||
from PIL import Image
|
||||
from typing import Dict, List, Optional, Tuple, Union, Any
|
||||
|
||||
from .utils import get_sd_version
|
||||
class easySampler:
|
||||
def __init__(self):
|
||||
self.last_helds: dict[str, list] = {
|
||||
"results": [],
|
||||
"pipe_line": [],
|
||||
}
|
||||
self.device = comfy.model_management.intermediate_device()
|
||||
|
||||
@staticmethod
|
||||
def tensor2pil(image: torch.Tensor) -> Image.Image:
|
||||
@@ -47,18 +48,41 @@ class easySampler:
|
||||
parts.append('None')
|
||||
return parts
|
||||
|
||||
def emptyLatent(self, resolution, empty_latent_width, empty_latent_height, batch_size=1, compression=0):
|
||||
if resolution != "自定义 x 自定义":
|
||||
try:
|
||||
width, height = map(int, resolution.split(' x '))
|
||||
empty_latent_width = width
|
||||
empty_latent_height = height
|
||||
except ValueError:
|
||||
raise ValueError("Invalid base_resolution format.")
|
||||
|
||||
if compression == 0:
|
||||
latent = torch.zeros([batch_size, 4, empty_latent_height // 8, empty_latent_width // 8], device=self.device)
|
||||
samples = {"samples": latent}
|
||||
else:
|
||||
latent_c = torch.zeros(
|
||||
[batch_size, 16, empty_latent_height // compression, empty_latent_width // compression])
|
||||
latent_b = torch.zeros([batch_size, 4, empty_latent_height // 4, empty_latent_width // 4])
|
||||
|
||||
samples = ({"samples": latent_c}, {"samples": latent_b})
|
||||
return samples
|
||||
|
||||
def add_model_patch_option(self, model):
|
||||
if 'transformer_options' not in model.model_options:
|
||||
model.model_options['transformer_options'] = {}
|
||||
to = model.model_options['transformer_options']
|
||||
if "model_patch" not in to:
|
||||
to["model_patch"] = {}
|
||||
return to
|
||||
|
||||
|
||||
def common_ksampler(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent, denoise=1.0,
|
||||
disable_noise=False, start_step=None, last_step=None, force_full_denoise=False,
|
||||
preview_latent=True, disable_pbar=False):
|
||||
device = comfy.model_management.get_torch_device()
|
||||
latent_image = latent["samples"]
|
||||
|
||||
if disable_noise:
|
||||
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
|
||||
else:
|
||||
batch_inds = latent["batch_index"] if "batch_index" in latent else None
|
||||
noise = comfy.sample.prepare_noise(latent_image, seed, batch_inds)
|
||||
|
||||
noise_mask = None
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent["noise_mask"]
|
||||
@@ -80,6 +104,34 @@ class easySampler:
|
||||
preview_bytes = previewer.decode_latent_to_preview_image(preview_format, x0)
|
||||
pbar.update_absolute(step + 1, total_steps, preview_bytes)
|
||||
|
||||
|
||||
if disable_noise:
|
||||
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout,
|
||||
device="cpu")
|
||||
else:
|
||||
batch_inds = latent["batch_index"] if "batch_index" in latent else None
|
||||
noise = comfy.sample.prepare_noise(latent_image, seed, batch_inds)
|
||||
|
||||
#######################################################################################
|
||||
# brushnet
|
||||
transformer_options = model.model_options['transformer_options'] if "transformer_options" in model.model_options else {}
|
||||
if 'model_patch' in transformer_options and 'brushnet' in transformer_options['model_patch']:
|
||||
to = self.add_model_patch_option(model)
|
||||
mp = to['model_patch']
|
||||
if isinstance(model.model.model_config, comfy.supported_models.SD15):
|
||||
mp['SDXL'] = False
|
||||
elif isinstance(model.model.model_config, comfy.supported_models.SDXL):
|
||||
mp['SDXL'] = True
|
||||
else:
|
||||
print('Base model type: ', type(model.model.model_config))
|
||||
raise Exception("Unsupported model type: ", type(model.model.model_config))
|
||||
|
||||
mp['unet'] = model.model.diffusion_model
|
||||
mp['step'] = 0
|
||||
mp['total_steps'] = 1
|
||||
|
||||
#
|
||||
#######################################################################################
|
||||
samples = comfy.sample.sample(model, noise, steps, cfg, sampler_name, scheduler, positive, negative,
|
||||
latent_image,
|
||||
denoise=denoise, disable_noise=disable_noise, start_step=start_step,
|
||||
@@ -118,8 +170,31 @@ class easySampler:
|
||||
|
||||
pbar = comfy.utils.ProgressBar(steps)
|
||||
|
||||
#######################################################################################
|
||||
# brushnet
|
||||
to = None
|
||||
transformer_options = model.model_options['transformer_options'] if "transformer_options" in model.model_options else {}
|
||||
if 'model_patch' in transformer_options and 'brushnet' in transformer_options['model_patch']:
|
||||
to = self.add_model_patch_option(model)
|
||||
mp = to['model_patch']
|
||||
if isinstance(model.model.model_config, comfy.supported_models.SD15):
|
||||
mp['SDXL'] = False
|
||||
elif isinstance(model.model.model_config, comfy.supported_models.SDXL):
|
||||
mp['SDXL'] = True
|
||||
else:
|
||||
print('Base model type: ', type(model.model.model_config))
|
||||
raise Exception("Unsupported model type: ", type(model.model.model_config))
|
||||
|
||||
mp['unet'] = model.model.diffusion_model
|
||||
mp['step'] = 0
|
||||
mp['total_steps'] = 1
|
||||
#
|
||||
#######################################################################################
|
||||
|
||||
def callback(step, x0, x, total_steps):
|
||||
preview_bytes = None
|
||||
if to is not None and "model_patch" in to:
|
||||
to['model_patch']['step'] = step + 1
|
||||
if previewer:
|
||||
preview_bytes = previewer.decode_latent_to_preview_image(preview_format, x0)
|
||||
pbar.update_absolute(step + 1, total_steps, preview_bytes)
|
||||
@@ -132,6 +207,32 @@ class easySampler:
|
||||
out["samples"] = samples
|
||||
return out
|
||||
|
||||
def custom_advanced_ksampler(self, noise, guider, sampler, sigmas, latent_image):
|
||||
latent = latent_image
|
||||
latent_image = latent["samples"]
|
||||
|
||||
noise_mask = None
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent["noise_mask"]
|
||||
|
||||
x0_output = {}
|
||||
callback = latent_preview.prepare_callback(guider.model_patcher, sigmas.shape[-1] - 1, x0_output)
|
||||
|
||||
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
|
||||
samples = guider.sample(noise.generate_noise(latent), latent_image, sampler, sigmas, denoise_mask=noise_mask,
|
||||
callback=callback, disable_pbar=disable_pbar, seed=noise.seed)
|
||||
samples = samples.to(comfy.model_management.intermediate_device())
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = samples
|
||||
if "x0" in x0_output:
|
||||
out_denoised = latent.copy()
|
||||
out_denoised["samples"] = guider.model_patcher.model.process_latent_out(x0_output["x0"].cpu())
|
||||
else:
|
||||
out_denoised = out
|
||||
|
||||
return (out, out_denoised)
|
||||
|
||||
def get_value_by_id(self, key: str, my_unique_id: Any) -> Optional[Any]:
|
||||
"""Retrieve value by its associated ID."""
|
||||
try:
|
||||
@@ -209,4 +310,43 @@ class easySampler:
|
||||
sdxl_pipe.get("clip"),
|
||||
sdxl_pipe.get("images"),
|
||||
sdxl_pipe.get("seed")
|
||||
)
|
||||
)
|
||||
|
||||
class alignYourStepsScheduler:
|
||||
|
||||
NOISE_LEVELS = {
|
||||
"SD1": [14.6146412293, 6.4745760956, 3.8636745985, 2.6946151520, 1.8841921177, 1.3943805092, 0.9642583904,
|
||||
0.6523686016, 0.3977456272, 0.1515232662, 0.0291671582],
|
||||
"SDXL": [14.6146412293, 6.3184485287, 3.7681790315, 2.1811480769, 1.3405244945, 0.8620721141, 0.5550693289,
|
||||
0.3798540708, 0.2332364134, 0.1114188177, 0.0291671582],
|
||||
"SVD": [700.00, 54.5, 15.886, 7.977, 4.248, 1.789, 0.981, 0.403, 0.173, 0.034, 0.002]}
|
||||
|
||||
|
||||
def loglinear_interp(self, t_steps, num_steps):
|
||||
"""
|
||||
Performs log-linear interpolation of a given array of decreasing numbers.
|
||||
"""
|
||||
xs = np.linspace(0, 1, len(t_steps))
|
||||
ys = np.log(t_steps[::-1])
|
||||
|
||||
new_xs = np.linspace(0, 1, num_steps)
|
||||
new_ys = np.interp(new_xs, xs, ys)
|
||||
|
||||
interped_ys = np.exp(new_ys)[::-1].copy()
|
||||
return interped_ys
|
||||
|
||||
def get_sigmas(self, model_type, steps, denoise):
|
||||
|
||||
total_steps = steps
|
||||
if denoise < 1.0:
|
||||
if denoise <= 0.0:
|
||||
return (torch.FloatTensor([]),)
|
||||
total_steps = round(steps * denoise)
|
||||
|
||||
sigmas = self.NOISE_LEVELS[model_type][:]
|
||||
if (steps + 1) != len(sigmas):
|
||||
sigmas = self.loglinear_interp(sigmas, steps + 1)
|
||||
|
||||
sigmas = sigmas[-(total_steps + 1):]
|
||||
sigmas[-1] = 0
|
||||
return (torch.FloatTensor(sigmas),)
|
||||
@@ -0,0 +1,201 @@
|
||||
import json
|
||||
import os
|
||||
import yaml
|
||||
import requests
|
||||
import pathlib
|
||||
from aiohttp import web
|
||||
from server import PromptServer
|
||||
from .image import tensor2pil, pil2tensor, image2base64, pil2byte
|
||||
from .log import log_node_error
|
||||
|
||||
|
||||
root_path = pathlib.Path(__file__).parent.parent.parent
|
||||
config_path = os.path.join(root_path,'config.yaml')
|
||||
default_key = [{'name':'Default', 'key':''}]
|
||||
|
||||
class StabilityAPI:
|
||||
def __init__(self):
|
||||
self.api_url = "https://api.stability.ai"
|
||||
self.api_keys = None
|
||||
self.api_current = 0
|
||||
self.user_info = {}
|
||||
self.getAPIKeys()
|
||||
|
||||
def getErrors(self, code):
|
||||
errors = {
|
||||
400: "Bad Request",
|
||||
403: "ApiKey Forbidden",
|
||||
413: "Your request was larger than 10MiB.",
|
||||
429: "You have made more than 150 requests in 10 seconds.",
|
||||
500: "Internal Server Error",
|
||||
}
|
||||
return errors.get(code, "Unknown Error")
|
||||
|
||||
def getAPIKeys(self):
|
||||
if os.path.isfile(config_path):
|
||||
with open(config_path, 'r') as f:
|
||||
data = yaml.load(f, Loader=yaml.FullLoader)
|
||||
if not data:
|
||||
data = {'STABILITY_API_KEY': default_key, 'STABILITY_API_DEFAULT':0}
|
||||
with open(config_path, 'w') as f:
|
||||
yaml.dump(data, f)
|
||||
if 'STABILITY_API_KEY' not in data:
|
||||
data['STABILITY_API_KEY'] = default_key
|
||||
data['STABILITY_API_DEFAULT'] = 0
|
||||
with open(config_path, 'w') as f:
|
||||
yaml.dump(data, f)
|
||||
api_keys = data['STABILITY_API_KEY']
|
||||
self.api_current = data['STABILITY_API_DEFAULT']
|
||||
self.api_keys = api_keys
|
||||
return api_keys
|
||||
else:
|
||||
# create a yaml file
|
||||
with open(config_path, 'w') as f:
|
||||
data = {'STABILITY_API_KEY': default_key, 'STABILITY_API_DEFAULT':0}
|
||||
yaml.dump(data, f)
|
||||
return data['STABILITY_API_KEY']
|
||||
pass
|
||||
|
||||
def setAPIKeys(self, api_keys):
|
||||
if len(api_keys) > 0:
|
||||
self.api_keys = api_keys
|
||||
# load and save the yaml file
|
||||
with open(config_path, 'r') as f:
|
||||
data = yaml.load(f, Loader=yaml.FullLoader)
|
||||
data['STABILITY_API_KEY'] = api_keys
|
||||
with open(config_path, 'w') as f:
|
||||
yaml.dump(data, f)
|
||||
return True
|
||||
|
||||
def setAPIDefault(self, current):
|
||||
if current is not None:
|
||||
self.api_current = current
|
||||
# load and save the yaml file
|
||||
with open(config_path, 'r') as f:
|
||||
data = yaml.load(f, Loader=yaml.FullLoader)
|
||||
data['STABILITY_API_DEFAULT'] = current
|
||||
with open(config_path, 'w') as f:
|
||||
yaml.dump(data, f)
|
||||
return True
|
||||
|
||||
def generate_sd3_image(self, prompt, negative_prompt, aspect_ratio, model, seed, mode='text-to-image', image=None, strength=1, output_format='png', node_name='easy stableDiffusion3API'):
|
||||
url = f"{self.api_url}/v2beta/stable-image/generate/sd3"
|
||||
api_key = self.api_keys[self.api_current]['key']
|
||||
files = None
|
||||
data = {
|
||||
"prompt": prompt,
|
||||
"mode": mode,
|
||||
"model": model,
|
||||
"seed": seed,
|
||||
"output_format": output_format,
|
||||
}
|
||||
if model == 'sd3':
|
||||
data['negative_prompt'] = negative_prompt
|
||||
|
||||
if mode == 'text-to-image':
|
||||
files = {"none": ''}
|
||||
data['aspect_ratio'] = aspect_ratio
|
||||
elif mode == 'image-to-image':
|
||||
pil_image = tensor2pil(image)
|
||||
image_byte = pil2byte(pil_image)
|
||||
files = {"image": ("output.png", image_byte, 'image/png')}
|
||||
data['strength'] = strength
|
||||
|
||||
response = requests.post(url,
|
||||
headers={"authorization": f"{api_key}", "accept": "application/json"},
|
||||
files=files,
|
||||
data=data,
|
||||
)
|
||||
if response.status_code == 200:
|
||||
PromptServer.instance.send_sync('stable-diffusion-api-generate-succeed',{"model":model})
|
||||
json_data = response.json()
|
||||
image_base64 = json_data['image']
|
||||
image_data = image2base64(image_base64)
|
||||
output_t = pil2tensor(image_data)
|
||||
return output_t
|
||||
else:
|
||||
if 'application/json' in response.headers['Content-Type']:
|
||||
error_info = response.json()
|
||||
log_node_error(node_name, error_info.get('name', 'No name provided'))
|
||||
log_node_error(node_name, error_info.get('errors', ['No details provided']))
|
||||
error_status_text = self.getErrors(response.status_code)
|
||||
PromptServer.instance.send_sync('easyuse-toast',{"type": "error", "content": error_status_text})
|
||||
raise Exception(f"Failed to generate image: {error_status_text}")
|
||||
|
||||
# get user account
|
||||
async def getUserAccount(self, cache=True):
|
||||
url = f"{self.api_url}/v1/user/account"
|
||||
api_key = self.api_keys[self.api_current]['key']
|
||||
name = self.api_keys[self.api_current]['name']
|
||||
if cache and name in self.user_info:
|
||||
return self.user_info[name]
|
||||
else:
|
||||
response = requests.get(url, headers={"Authorization": f"Bearer {api_key}"})
|
||||
if response.status_code == 200:
|
||||
user_info = response.json()
|
||||
self.user_info[name] = user_info
|
||||
return user_info
|
||||
else:
|
||||
PromptServer.instance.send_sync('easyuse-toast',{'type': 'error', 'content': self.getErrors(response.status_code)})
|
||||
return None
|
||||
|
||||
# get user balance
|
||||
async def getUserBalance(self):
|
||||
url = f"{self.api_url}/v1/user/balance"
|
||||
api_key = self.api_keys[self.api_current]['key']
|
||||
response = requests.get(url, headers={
|
||||
"Authorization": f"Bearer {api_key}"
|
||||
})
|
||||
if response.status_code == 200:
|
||||
return response.json()
|
||||
else:
|
||||
PromptServer.instance.send_sync('easyuse-toast', {'type': 'error', 'content': self.getErrors(response.status_code)})
|
||||
return None
|
||||
|
||||
stableAPI = StabilityAPI()
|
||||
|
||||
|
||||
@PromptServer.instance.routes.get("/easyuse/stability/api_keys")
|
||||
async def get_stability_api_keys(request):
|
||||
stableAPI.getAPIKeys()
|
||||
return web.json_response({"keys": stableAPI.api_keys, "current": stableAPI.api_current})
|
||||
|
||||
@PromptServer.instance.routes.post("/easyuse/stability/set_api_keys")
|
||||
async def set_stability_api_keys(request):
|
||||
post = await request.post()
|
||||
api_keys = post.get("api_keys")
|
||||
current = post.get('current')
|
||||
if api_keys is not None:
|
||||
api_keys = json.loads(api_keys)
|
||||
stableAPI.setAPIKeys(api_keys)
|
||||
if current is not None:
|
||||
print(current)
|
||||
stableAPI.setAPIDefault(int(current))
|
||||
account = await stableAPI.getUserAccount()
|
||||
balance = await stableAPI.getUserBalance()
|
||||
return web.json_response({'account': account, 'balance': balance})
|
||||
else:
|
||||
return web.json_response({'status': 'ok'})
|
||||
else:
|
||||
return web.Response(status=400)
|
||||
|
||||
@PromptServer.instance.routes.post("/easyuse/stability/set_apikey_default")
|
||||
async def set_stability_api_default(request):
|
||||
post = await request.post()
|
||||
current = post.get("current")
|
||||
if current is not None and current < len(stableAPI.api_keys):
|
||||
stableAPI.api_current = current
|
||||
return web.json_response({'status': 'ok'})
|
||||
else:
|
||||
return web.Response(status=400)
|
||||
|
||||
@PromptServer.instance.routes.get("/easyuse/stability/user_info")
|
||||
async def get_account_info(request):
|
||||
account = await stableAPI.getUserAccount()
|
||||
balance = await stableAPI.getUserBalance()
|
||||
return web.json_response({'account': account, 'balance': balance})
|
||||
|
||||
@PromptServer.instance.routes.get("/easyuse/stability/balance")
|
||||
async def get_balance_info(request):
|
||||
balance = await stableAPI.getUserBalance()
|
||||
return web.json_response({'balance': balance})
|
||||
@@ -0,0 +1,148 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from typing import Union
|
||||
|
||||
T = torch.Tensor
|
||||
|
||||
|
||||
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"]
|
||||
|
||||
|
||||
def styleAlignBatch(model, share_norm, share_attn, scale=1.0):
|
||||
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
|
||||
@@ -0,0 +1,238 @@
|
||||
import re
|
||||
import os
|
||||
import folder_paths
|
||||
|
||||
import comfy.utils
|
||||
import torch
|
||||
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
|
||||
|
||||
from .utils import install_package
|
||||
try:
|
||||
from lark import Lark, Transformer, v_args
|
||||
except:
|
||||
print('install lark-parser...')
|
||||
install_package('lark-parser')
|
||||
from lark import Lark, Transformer, v_args
|
||||
|
||||
model_path = os.path.join(folder_paths.models_dir, 'prompt_generator')
|
||||
zh_en_model_path = os.path.join(model_path, 'opus-mt-zh-en')
|
||||
zh_en_model, zh_en_tokenizer = None, None
|
||||
|
||||
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)
|
||||
|
||||
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"
|
||||
|
||||
def has_chinese(text):
|
||||
has_cn = False
|
||||
_text = text
|
||||
_text = re.sub(r'<.*?>', '', _text)
|
||||
_text = re.sub(r'__.*?__', '', _text)
|
||||
_text = re.sub(r'embedding:.*?(\d+)?', '', _text)
|
||||
for char in _text:
|
||||
if '\u4e00' <= char <= '\u9fff':
|
||||
has_cn = True
|
||||
break
|
||||
elif char.isalpha():
|
||||
continue
|
||||
return has_cn
|
||||
|
||||
def translate(text):
|
||||
global zh_en_model_path, zh_en_model, zh_en_tokenizer
|
||||
|
||||
if not os.path.exists(zh_en_model_path):
|
||||
zh_en_model_path = 'Helsinki-NLP/opus-mt-zh-en'
|
||||
|
||||
if zh_en_model is 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]
|
||||
|
||||
@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):
|
||||
if len(args) == 1:
|
||||
return f"<lora:{args[0]}>"
|
||||
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 re.search(r'__.*?__', str(word)):
|
||||
return str(word).rstrip('.')
|
||||
elif re.search(r'@.*?@', str(word)):
|
||||
return str(word).replace('@', '').rstrip('.')
|
||||
elif detect_language(str(word)) == "cn":
|
||||
return translate(str(word)).rstrip('.')
|
||||
else:
|
||||
return str(word).rstrip('.')
|
||||
|
||||
|
||||
#定义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: /[^,:\(\)\[\]<>]+/
|
||||
"""
|
||||
def zh_to_en(text):
|
||||
global zh_en_model_path, zh_en_model, zh_en_tokenizer
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(len(text) + 1)
|
||||
texts = [correct_prompt_syntax(t) for t in text]
|
||||
|
||||
install_package('sentencepiece', '0.2.0')
|
||||
|
||||
if not os.path.exists(zh_en_model_path):
|
||||
zh_en_model_path = 'Helsinki-NLP/opus-mt-zh-en'
|
||||
|
||||
if zh_en_model is 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")
|
||||
|
||||
prompt_result = []
|
||||
|
||||
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:
|
||||
prompt_result.append(t)
|
||||
pbar.update(1)
|
||||
|
||||
# print('prompt_result', prompt_result, )
|
||||
if len(prompt_result) == 0:
|
||||
prompt_result = [""]
|
||||
|
||||
return prompt_result
|
||||
+112
-40
@@ -1,3 +1,10 @@
|
||||
class AlwaysEqualProxy(str):
|
||||
def __eq__(self, _):
|
||||
return True
|
||||
|
||||
def __ne__(self, _):
|
||||
return False
|
||||
|
||||
comfy_ui_revision = None
|
||||
def get_comfyui_revision():
|
||||
try:
|
||||
@@ -10,22 +17,72 @@ def get_comfyui_revision():
|
||||
comfy_ui_revision = "Unknown"
|
||||
return comfy_ui_revision
|
||||
|
||||
|
||||
import sys
|
||||
import importlib.util
|
||||
import importlib.metadata
|
||||
import comfy.model_management as mm
|
||||
import gc
|
||||
from packaging import version
|
||||
from server import PromptServer
|
||||
def is_package_installed(package):
|
||||
try:
|
||||
module = importlib.util.find_spec(package)
|
||||
return module is not None
|
||||
except ImportError as e:
|
||||
print(e)
|
||||
return False
|
||||
|
||||
def install_package(package, v=None, compare=True, compare_version=None):
|
||||
run_install = True
|
||||
if is_package_installed(package):
|
||||
try:
|
||||
installed_version = importlib.metadata.version(package)
|
||||
if v is not None:
|
||||
if compare_version is None:
|
||||
compare_version = v
|
||||
if not compare or version.parse(installed_version) >= version.parse(compare_version):
|
||||
run_install = False
|
||||
else:
|
||||
run_install = False
|
||||
except:
|
||||
run_install = False
|
||||
|
||||
if run_install:
|
||||
import subprocess
|
||||
package_command = package + '==' + v if v is not None else package
|
||||
PromptServer.instance.send_sync("easyuse-toast", {'content': f"Installing {package_command}...", 'duration': 5000})
|
||||
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', package_command], capture_output=True, text=True)
|
||||
if result.returncode == 0:
|
||||
PromptServer.instance.send_sync("easyuse-toast", {'content': f"{package} installed successfully", 'type': 'success', 'duration': 5000})
|
||||
print(f"Package {package} installed successfully")
|
||||
return True
|
||||
else:
|
||||
PromptServer.instance.send_sync("easyuse-toast", {'content': f"{package} installed failed", 'type': 'error', 'duration': 5000})
|
||||
print(f"Package {package} installed failed")
|
||||
return False
|
||||
else:
|
||||
return False
|
||||
|
||||
def compare_revision(num):
|
||||
global comfy_ui_revision
|
||||
if not comfy_ui_revision:
|
||||
comfy_ui_revision = get_comfyui_revision()
|
||||
return True if comfy_ui_revision == 'Unknown' or int(comfy_ui_revision) >= num else False
|
||||
def find_tags(string: str, sep="/") -> list[str]:
|
||||
"""
|
||||
find tags from string use the sep for split
|
||||
Note: string may contain the \\ or / for path separator
|
||||
"""
|
||||
if not string:
|
||||
return []
|
||||
string = string.replace("\\", "/")
|
||||
while "//" in string:
|
||||
string = string.replace("//", "/")
|
||||
if string and sep in string:
|
||||
return string.split(sep)[:-1]
|
||||
return []
|
||||
|
||||
import folder_paths
|
||||
def add_folder_path_and_extensions(folder_name, full_folder_paths, extensions):
|
||||
for full_folder_path in full_folder_paths:
|
||||
folder_paths.add_model_folder_path(folder_name, full_folder_path)
|
||||
if folder_name in folder_paths.folder_names_and_paths:
|
||||
current_paths, current_extensions = folder_paths.folder_names_and_paths[folder_name]
|
||||
updated_extensions = current_extensions | extensions
|
||||
folder_paths.folder_names_and_paths[folder_name] = (current_paths, updated_extensions)
|
||||
else:
|
||||
folder_paths.folder_names_and_paths[folder_name] = (full_folder_paths, extensions)
|
||||
|
||||
from comfy.model_base import BaseModel
|
||||
import comfy.supported_models
|
||||
@@ -38,7 +95,11 @@ def get_sd_version(model):
|
||||
elif isinstance(
|
||||
model_config, (comfy.supported_models.SD15, comfy.supported_models.SD20)
|
||||
):
|
||||
return 'sd15'
|
||||
return 'sd1'
|
||||
elif isinstance(
|
||||
model_config, (comfy.supported_models.SVD_img2vid)
|
||||
):
|
||||
return 'svd'
|
||||
else:
|
||||
return 'unknown'
|
||||
|
||||
@@ -108,11 +169,14 @@ def is_linked_styles_selector(prompt, my_unique_id, prompt_type='positive'):
|
||||
else:
|
||||
return False
|
||||
|
||||
use_mirror = False
|
||||
def get_local_filepath(url, dirname, local_file_name=None):
|
||||
"""Get local file path when is already downloaded or download it"""
|
||||
import os
|
||||
from server import PromptServer
|
||||
from urllib.parse import urlparse
|
||||
from torch.hub import download_url_to_file
|
||||
global use_mirror
|
||||
if not os.path.exists(dirname):
|
||||
os.makedirs(dirname)
|
||||
if not local_file_name:
|
||||
@@ -120,8 +184,23 @@ def get_local_filepath(url, dirname, local_file_name=None):
|
||||
local_file_name = os.path.basename(parsed_url.path)
|
||||
destination = os.path.join(dirname, local_file_name)
|
||||
if not os.path.exists(destination):
|
||||
print(f'downloading {url} to {destination}')
|
||||
download_url_to_file(url, destination)
|
||||
try:
|
||||
if use_mirror:
|
||||
url = url.replace('huggingface.co', 'hf-mirror.com')
|
||||
print(f'downloading {url} to {destination}')
|
||||
PromptServer.instance.send_sync("easyuse-toast", {'content': f'Downloading model to {destination}, please wait...', 'duration': 10000})
|
||||
download_url_to_file(url, destination)
|
||||
except Exception as e:
|
||||
use_mirror = True
|
||||
url = url.replace('huggingface.co', 'hf-mirror.com')
|
||||
print(f'无法从huggingface下载,正在尝试从 {url} 下载...')
|
||||
PromptServer.instance.send_sync("easyuse-toast", {'content': f'无法连接huggingface,正在尝试从 {url} 下载...', 'duration': 10000})
|
||||
try:
|
||||
download_url_to_file(url, destination)
|
||||
except Exception as err:
|
||||
PromptServer.instance.send_sync("easyuse-toast",
|
||||
{'content': f'无法从 {url} 下载模型', 'type':'error'})
|
||||
raise Exception(f'无法从 {url} 下载,错误信息:{str(err.args[0])}')
|
||||
return destination
|
||||
|
||||
def to_lora_patch_dict(state_dict: dict) -> dict:
|
||||
@@ -148,7 +227,7 @@ def easySave(images, filename_prefix, output_type, prompt=None, extra_pnginfo=No
|
||||
from nodes import PreviewImage, SaveImage
|
||||
if output_type == "Hide":
|
||||
return list()
|
||||
if output_type == "Preview":
|
||||
if output_type in ["Preview", "Preview&Choose"]:
|
||||
filename_prefix = 'easyPreview'
|
||||
results = PreviewImage().save_images(images, filename_prefix, prompt, extra_pnginfo)
|
||||
return results['ui']['images']
|
||||
@@ -156,29 +235,22 @@ def easySave(images, filename_prefix, output_type, prompt=None, extra_pnginfo=No
|
||||
results = SaveImage().save_images(images, filename_prefix, prompt, extra_pnginfo)
|
||||
return results['ui']['images']
|
||||
|
||||
# Image Utils
|
||||
# from PIL import Image, ImageDraw
|
||||
# import numpy as np
|
||||
# import torch
|
||||
# def is_image_transparent(img):
|
||||
# print(img.shape)
|
||||
# if len(img.shape) > 3 and img.shape[3] == 4:
|
||||
# return True
|
||||
# else:
|
||||
# m = tensor2pil(img)
|
||||
# if m.mode == "RGBA":
|
||||
# return True
|
||||
# else:
|
||||
# return False
|
||||
#
|
||||
# def create_grid(image_size, box_size):
|
||||
# img = Image.new('RGBA', image_size, (255, 255, 255, 255)) # 白色背景
|
||||
# draw = ImageDraw.Draw(img)
|
||||
#
|
||||
# for x in range(0, img.width, box_size):
|
||||
# for y in range(0, img.height, box_size):
|
||||
# if (x // box_size % 2 == 0 and y // box_size % 2 == 0) or (x // box_size % 2 == 1 and y // box_size % 2 == 1):
|
||||
# draw.rectangle([(x, y), (x+box_size, y+box_size)], fill=(204, 204, 204, 255)) # 不透明
|
||||
# else:
|
||||
# continue # 保持透明
|
||||
# return img
|
||||
def getMetadata(filepath):
|
||||
with open(filepath, "rb") as file:
|
||||
# https://github.com/huggingface/safetensors#format
|
||||
# 8 bytes: N, an unsigned little-endian 64-bit integer, containing the size of the header
|
||||
header_size = int.from_bytes(file.read(8), "little", signed=False)
|
||||
|
||||
if header_size <= 0:
|
||||
raise BufferError("Invalid header size")
|
||||
|
||||
header = file.read(header_size)
|
||||
if header_size <= 0:
|
||||
raise BufferError("Invalid header")
|
||||
|
||||
return header
|
||||
|
||||
def cleanGPUUsedForce():
|
||||
gc.collect()
|
||||
mm.unload_all_models()
|
||||
mm.soft_empty_cache()
|
||||
@@ -156,9 +156,8 @@ def process(text, seed=None):
|
||||
|
||||
def replace_wildcard(string):
|
||||
global easy_wildcard_dict
|
||||
pattern = r"__([\w.\-+/*\\]+)__"
|
||||
pattern = r"__([\w\s.\-+/*\\]+?)__"
|
||||
matches = re.findall(pattern, string)
|
||||
|
||||
replacements_found = False
|
||||
|
||||
for match in matches:
|
||||
+25
-7
@@ -2,10 +2,12 @@ import os, torch
|
||||
from pathlib import Path
|
||||
from PIL import Image, ImageDraw, ImageFont
|
||||
from .utils import easySave
|
||||
from ..config import RESOURCES_DIR
|
||||
from ..log import log_node_warn
|
||||
from ..adv_encode import advanced_encode
|
||||
from .adv_encode import advanced_encode
|
||||
from .controlnet import easyControlnet
|
||||
from .log import log_node_warn
|
||||
from ..layer_diffuse.func import LayerDiffuse
|
||||
from ..config import RESOURCES_DIR
|
||||
|
||||
class easyXYPlot():
|
||||
|
||||
def __init__(self, xyPlotData, save_prefix, image_output, prompt, extra_pnginfo, my_unique_id, sampler, easyCache):
|
||||
@@ -307,6 +309,7 @@ class easyXYPlot():
|
||||
ckpt_name, clip_skip, vae_name = xy_values.split(",")
|
||||
ckpt_name = ckpt_name.replace('*', ',')
|
||||
vae_name = vae_name.replace('*', ',')
|
||||
print(ckpt_name)
|
||||
model, clip, vae, clip_vision = self.easyCache.load_checkpoint(ckpt_name)
|
||||
if vae_name != 'None':
|
||||
vae = self.easyCache.load_vae(vae_name)
|
||||
@@ -315,6 +318,8 @@ class easyXYPlot():
|
||||
optional_lora_stack = plot_image_vars['lora_stack']
|
||||
if optional_lora_stack is not None and optional_lora_stack != []:
|
||||
for lora in optional_lora_stack:
|
||||
lora['model'] = model
|
||||
lora['clip'] = clip
|
||||
model, clip = self.easyCache.load_lora(lora)
|
||||
|
||||
# 处理clip
|
||||
@@ -354,8 +359,9 @@ class easyXYPlot():
|
||||
plot_image_vars['positive_weight_interpretation'],
|
||||
w_max=1.0,
|
||||
apply_to_pooled="enable", a1111_prompt_style=a1111_prompt_style, steps=steps)
|
||||
if "positive_cond" in plot_image_vars:
|
||||
positive = positive + plot_image_vars["positive_cond"]
|
||||
|
||||
# if "positive_cond" in plot_image_vars:
|
||||
# positive = positive + plot_image_vars["positive_cond"]
|
||||
|
||||
if "Negative" in self.x_type or "Negative" in self.y_type:
|
||||
if self.x_type == 'Negative Prompt S/R' or self.y_type == 'Negative Prompt S/R':
|
||||
@@ -366,8 +372,8 @@ class easyXYPlot():
|
||||
plot_image_vars['negative_weight_interpretation'],
|
||||
w_max=1.0,
|
||||
apply_to_pooled="enable", a1111_prompt_style=a1111_prompt_style, steps=steps)
|
||||
if "negative_cond" in plot_image_vars:
|
||||
negative = negative + plot_image_vars["negative_cond"]
|
||||
# if "negative_cond" in plot_image_vars:
|
||||
# negative = negative + plot_image_vars["negative_cond"]
|
||||
|
||||
# ControlNet
|
||||
if "ControlNet" in self.x_type or "ControlNet" in self.y_type:
|
||||
@@ -428,8 +434,20 @@ class easyXYPlot():
|
||||
# LayerDiffuse
|
||||
layer_diffusion_method = plot_image_vars["layer_diffusion_method"] if "layer_diffusion_method" in plot_image_vars else None
|
||||
empty_samples = plot_image_vars["empty_samples"] if "empty_samples" in plot_image_vars else None
|
||||
|
||||
if layer_diffusion_method:
|
||||
samp_blend_samples = plot_image_vars["blend_samples"] if "blend_samples" in plot_image_vars else None
|
||||
additional_cond = plot_image_vars["layer_diffusion_cond"] if "layer_diffusion_cond" in plot_image_vars else None
|
||||
|
||||
images = plot_image_vars["images"].movedim(-1, 1) if "images" in plot_image_vars else None
|
||||
weight = plot_image_vars['layer_diffusion_weight'] if 'layer_diffusion_weight' in plot_image_vars else 1.0
|
||||
model, positive, negative = LayerDiffuse().apply_layer_diffusion(model, layer_diffusion_method, weight, samples,
|
||||
samp_blend_samples, positive,
|
||||
negative, images, additional_cond)
|
||||
|
||||
samples = empty_samples if layer_diffusion_method is not None and empty_samples is not None else samples
|
||||
# Sample
|
||||
|
||||
samples = self.sampler.common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, samples,
|
||||
denoise=denoise, disable_noise=disable_noise, preview_latent=preview_latent,
|
||||
start_step=start_step, last_step=last_step,
|
||||
|
||||
+77
-11
@@ -1,7 +1,8 @@
|
||||
from typing import Iterator, List, Tuple, Dict, Any, Union, Optional
|
||||
from _decimal import Context, getcontext
|
||||
from decimal import Decimal
|
||||
import torch
|
||||
from .libs.utils import AlwaysEqualProxy, cleanGPUUsedForce
|
||||
from .libs.cache import remove_cache
|
||||
import numpy as np
|
||||
import json
|
||||
|
||||
@@ -277,6 +278,30 @@ class imageSwitch:
|
||||
else:
|
||||
return (image_b, )
|
||||
|
||||
class textSwitch:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"input": ("INT", {"default": 1, "min": 1, "max": 2}),
|
||||
},
|
||||
"optional": {
|
||||
"text1": ("STRING", {"forceInput": True}),
|
||||
"text2": ("STRING", {"forceInput": True}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("STRING",)
|
||||
CATEGORY = "EasyUse/Logic/Switch"
|
||||
FUNCTION = "switch"
|
||||
|
||||
def switch(self, input, text1=None, text2=None,):
|
||||
if input == 1:
|
||||
return (text1,)
|
||||
else:
|
||||
return (text2,)
|
||||
|
||||
# ---------------------------------------------------------------运算 开始----------------------------------------------------------------------#
|
||||
|
||||
COMPARE_FUNCTIONS = {
|
||||
@@ -287,12 +312,6 @@ COMPARE_FUNCTIONS = {
|
||||
"a <= b": lambda a, b: a <= b,
|
||||
"a >= b": lambda a, b: a >= b,
|
||||
}
|
||||
class AlwaysEqualProxy(str):
|
||||
def __eq__(self, _):
|
||||
return True
|
||||
|
||||
def __ne__(self, _):
|
||||
return False
|
||||
|
||||
# 比较
|
||||
class Compare:
|
||||
@@ -512,11 +531,52 @@ class cleanGPUUsed:
|
||||
CATEGORY = "EasyUse/Logic"
|
||||
|
||||
def empty_cache(self, anything, unique_id=None, extra_pnginfo=None):
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
cleanGPUUsedForce()
|
||||
remove_cache('*')
|
||||
return ()
|
||||
|
||||
class clearCacheKey:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"anything": (AlwaysEqualProxy("*"), {}),
|
||||
"cache_key": ("STRING", {"default": "*"}),
|
||||
}, "optional": {},
|
||||
"hidden": {"unique_id": "UNIQUE_ID", "extra_pnginfo": "EXTRA_PNGINFO",}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
RETURN_NAMES = ()
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "empty_cache"
|
||||
CATEGORY = "EasyUse/Logic"
|
||||
|
||||
def empty_cache(self, anything, cache_name, unique_id=None, extra_pnginfo=None):
|
||||
remove_cache(cache_name)
|
||||
return ()
|
||||
|
||||
class clearCacheAll:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"anything": (AlwaysEqualProxy("*"), {}),
|
||||
}, "optional": {},
|
||||
"hidden": {"unique_id": "UNIQUE_ID", "extra_pnginfo": "EXTRA_PNGINFO",}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
RETURN_NAMES = ()
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "empty_cache"
|
||||
CATEGORY = "EasyUse/Logic"
|
||||
|
||||
def empty_cache(self, anything, unique_id=None, extra_pnginfo=None):
|
||||
remove_cache('*')
|
||||
return ()
|
||||
|
||||
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"easy string": String,
|
||||
"easy int": Int,
|
||||
@@ -526,13 +586,16 @@ NODE_CLASS_MAPPINGS = {
|
||||
"easy boolean": Boolean,
|
||||
"easy compare": Compare,
|
||||
"easy imageSwitch": imageSwitch,
|
||||
"easy textSwitch": textSwitch,
|
||||
"easy if": If,
|
||||
"easy isSDXL": isSDXL,
|
||||
"easy xyAny": xyAny,
|
||||
"easy convertAnything": ConvertAnything,
|
||||
"easy showAnything": showAnything,
|
||||
"easy showTensorShape": showTensorShape,
|
||||
"easy cleanGpuUsed": cleanGPUUsed
|
||||
"easy clearCacheKey": clearCacheKey,
|
||||
"easy clearCacheAll": clearCacheAll,
|
||||
"easy cleanGpuUsed": cleanGPUUsed,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"easy string": "String",
|
||||
@@ -543,11 +606,14 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"easy boolean": "Boolean",
|
||||
"easy compare": "Compare",
|
||||
"easy imageSwitch": "Image Switch",
|
||||
"easy textSwitch": "Text Switch",
|
||||
"easy if": "If",
|
||||
"easy isSDXL": "Is SDXL",
|
||||
"easy xyAny": "XYAny",
|
||||
"easy convertAnything": "Convert Any",
|
||||
"easy showAnything": "Show Any",
|
||||
"easy showTensorShape": "Show Tensor Shape",
|
||||
"easy clearCacheKey": "Clear Cache Key",
|
||||
"easy clearCacheAll": "Clear Cache All",
|
||||
"easy cleanGpuUsed": "Clean GPU Used"
|
||||
}
|
||||
@@ -2,9 +2,6 @@ import random
|
||||
import server
|
||||
from enum import Enum
|
||||
|
||||
|
||||
seed_nodes = ["easy wildcards","easy preSampling","easy preSamplingAdvanced","easy preSamplingSdTurbo","easy preSamplingDynamicCFG","easy preSamplingLayerDiffusion","easy preSamplingCascade","easy fullCascadeKSampler","easy fullkSampler","easy seed","easy latentNoisy", "easy preSamplingNoiseIn"]
|
||||
|
||||
class SGmode(Enum):
|
||||
FIX = 1
|
||||
INCR = 2
|
||||
@@ -123,41 +120,6 @@ def prompt_seed_update(json_data):
|
||||
# control after generated
|
||||
if mode is not None and not mode:
|
||||
control_seed(node[1], action, seed_is_global)
|
||||
# else:
|
||||
# prompts = json_data['prompt'].items()
|
||||
# for k, v in prompts:
|
||||
# if 'class_type' not in v:
|
||||
# continue
|
||||
# cls = v['class_type']
|
||||
# if cls in seed_nodes:
|
||||
# extra_data = next((x for x in workflow["nodes"] if str(x["id"]) == k), None)
|
||||
# if extra_data is not None:
|
||||
# inputs = extra_data.get('inputs')
|
||||
# widgets_value = extra_data.get('widgets_values')
|
||||
# widgets_length = len(widgets_value)
|
||||
# if "disable" in widgets_value:
|
||||
# break
|
||||
# if inputs is not None and inputs != []:
|
||||
# seed_num_input = next((x for x in inputs if x['name'] == 'seed_num' and x['type'] == 'INT'), None)
|
||||
# if seed_num_input is not None:
|
||||
# action = 'fixed'
|
||||
# else:
|
||||
# action = widgets_value[widgets_length - 1]
|
||||
# else:
|
||||
# control_index = widgets_length - 2 if cls == 'easy seed' else widgets_length - 1
|
||||
# action = widgets_value[control_index]
|
||||
#
|
||||
# # print(action)
|
||||
# node = k, v
|
||||
# value = control_seed(node[1], action, False)
|
||||
#
|
||||
# if k not in seed_widget_map:
|
||||
# continue
|
||||
#
|
||||
# if 'seed_num' in v['inputs']:
|
||||
# if isinstance(v['inputs']['seed_num'], int):
|
||||
# v['inputs']['seed_num'] = value
|
||||
|
||||
|
||||
return value is not None
|
||||
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
[project]
|
||||
name = "comfyui-easy-use"
|
||||
description = "To enhance the usability of ComfyUI, optimizations and integrations have been implemented for several commonly used nodes."
|
||||
version = "1.1.8"
|
||||
license = "LICENSE"
|
||||
dependencies = ["diffusers>=0.25.0", "clip_interrogator>=0.6.0", "sentencepiece==0.2.0", "lark-parser", "onnxruntime"]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/yolain/ComfyUI-Easy-Use"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "yolain"
|
||||
DisplayName = "ComfyUI-Easy-Use"
|
||||
Icon = ""
|
||||
+4
-1
@@ -1,2 +1,5 @@
|
||||
diffusers>=0.25.0
|
||||
aiohttp
|
||||
clip_interrogator>=0.6.0
|
||||
sentencepiece==0.2.0
|
||||
lark-parser
|
||||
onnxruntime
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
.easyuse-account{
|
||||
|
||||
}
|
||||
.easyuse-account-user{
|
||||
font-size: 10px;
|
||||
color:var(--descrip-text);
|
||||
text-align: center;
|
||||
}
|
||||
.easyuse-account-user-info{
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
padding-bottom:10px;
|
||||
cursor: pointer;
|
||||
}
|
||||
.easyuse-account-user-info .user{
|
||||
display: flex;
|
||||
align-items: center;
|
||||
}
|
||||
.easyuse-account-user-info .edit{
|
||||
padding:5px 10px;
|
||||
background: var(--comfy-menu-bg);
|
||||
border-radius:4px;
|
||||
}
|
||||
.easyuse-account-user-info:hover{
|
||||
filter:brightness(110%);
|
||||
}
|
||||
.easyuse-account-user-info h5{
|
||||
margin:0;
|
||||
font-size: 10px;
|
||||
text-align: left;
|
||||
}
|
||||
.easyuse-account-user-info h6{
|
||||
margin:0;
|
||||
font-size: 8px;
|
||||
text-align: left;
|
||||
font-weight: 300;
|
||||
}
|
||||
.easyuse-account-user-info .remark{
|
||||
margin-top: 4px;
|
||||
}
|
||||
.easyuse-account-user-info .avatar{
|
||||
width: 36px;
|
||||
height: 36px;
|
||||
background: var(--comfy-input-bg);
|
||||
border-radius: 50%;
|
||||
margin-right: 5px;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
font-size: 16px;
|
||||
overflow: hidden;
|
||||
}
|
||||
.easyuse-account-user-info .avatar img{
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
}
|
||||
.easyuse-account-dialog{
|
||||
width: 600px;
|
||||
}
|
||||
.easyuse-account-dialog-main a, .easyuse-account-dialog-main a:visited{
|
||||
font-weight: 400;
|
||||
color: var(--theme-color-light);
|
||||
}
|
||||
.easyuse-account-dialog-item{
|
||||
display: flex;
|
||||
justify-content: flex-start;
|
||||
align-items: center;
|
||||
padding: 10px 0;
|
||||
border-bottom: 1px solid var(--border-color);
|
||||
}
|
||||
.easyuse-account-dialog-item input{
|
||||
padding:5px;
|
||||
margin-right:5px;
|
||||
}
|
||||
.easyuse-account-dialog-item input.key{
|
||||
flex:1;
|
||||
}
|
||||
.easyuse-account-dialog-item button{
|
||||
cursor: pointer;
|
||||
margin-left:5px!important;
|
||||
padding:5px!important;
|
||||
font-size: 16px!important;
|
||||
}
|
||||
.easyuse-account-dialog-item button:hover{
|
||||
filter:brightness(120%);
|
||||
}
|
||||
.easyuse-account-dialog-item button.choose {
|
||||
background: var(--theme-color);
|
||||
}
|
||||
.easyuse-account-dialog-item button.delete{
|
||||
background: var(--error-color);
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
.easyuse-chooser-dialog{
|
||||
max-width: 600px;
|
||||
}
|
||||
.easyuse-chooser-dialog-title{
|
||||
font-size: 18px;
|
||||
font-weight: 700;
|
||||
text-align: center;
|
||||
color:var(--input-text);
|
||||
margin:0;
|
||||
}
|
||||
.easyuse-chooser-dialog-images{
|
||||
margin-top:10px;
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
width: 100%;
|
||||
box-sizing: border-box;
|
||||
}
|
||||
.easyuse-chooser-dialog-images img{
|
||||
width: 50%;
|
||||
height: auto;
|
||||
cursor: pointer;
|
||||
box-sizing: border-box;
|
||||
filter:brightness(80%);
|
||||
}
|
||||
.easyuse-chooser-dialog-images img:hover{
|
||||
filter:brightness(100%);
|
||||
}
|
||||
.easyuse-chooser-dialog-images img.selected{
|
||||
border: 4px solid var(--success-color);
|
||||
}
|
||||
|
||||
.easyuse-chooser-hidden{
|
||||
display: none;
|
||||
height:0;
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
.easyuse-model{
|
||||
position:relative;
|
||||
}
|
||||
.easyuse-model:hover img{
|
||||
display: block;
|
||||
opacity: 1;
|
||||
}
|
||||
.easyuse-model img{
|
||||
position: absolute;
|
||||
z-index:1;
|
||||
right:-155px;
|
||||
top:0;
|
||||
width:150px;
|
||||
height:auto;
|
||||
display: none;
|
||||
filter:brightness(70%);
|
||||
-webkit-filter: brightness(70%);
|
||||
opacity: 0;
|
||||
transition:all 0.5s ease-in-out;
|
||||
}
|
||||
@@ -10,4 +10,23 @@
|
||||
background-color: var(--comfy-menu-bg);
|
||||
padding: 10px 4px;
|
||||
border: 1px solid var(--border-color);z-index: 999999999;padding-top: 0;
|
||||
}
|
||||
#easyuse_groups_map .icon{
|
||||
width: 12px;
|
||||
height:12px;
|
||||
}
|
||||
#easyuse_groups_map .closeBtn{
|
||||
float: right;
|
||||
color: var(--input-text);
|
||||
border-radius:30px;
|
||||
background-color: var(--comfy-input-bg);
|
||||
border: 1px solid var(--border-color);
|
||||
cursor: pointer;
|
||||
aspect-ratio: 1 / 1;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
}
|
||||
#easyuse_groups_map .closeBtn:hover{
|
||||
filter:brightness(120%);
|
||||
}
|
||||
+7
-1
@@ -1,4 +1,10 @@
|
||||
@import "theme.css";
|
||||
@import "dropdown.css";
|
||||
@import "selector.css";
|
||||
@import "groupmap.css";
|
||||
@import "groupmap.css";
|
||||
@import "contextmenu.css";
|
||||
@import "modelinfo.css";
|
||||
@import "toast.css";
|
||||
@import "account.css";
|
||||
@import "chooser.css";
|
||||
@import "toolbar.css";
|
||||
@@ -0,0 +1,265 @@
|
||||
.easyuse-model-info {
|
||||
color: white;
|
||||
max-width: 90vw;
|
||||
font-family: var(--font-family);
|
||||
}
|
||||
.easyuse-model-content {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
overflow: hidden;
|
||||
}
|
||||
.easyuse-model-header{
|
||||
margin:0 0 15px 0;
|
||||
}
|
||||
.easyuse-model-header-remark{
|
||||
display: flex;
|
||||
align-items: center;
|
||||
margin-top:5px;
|
||||
}
|
||||
.easyuse-model-info h2 {
|
||||
text-align: left;
|
||||
margin:0;
|
||||
}
|
||||
.easyuse-model-info h5 {
|
||||
text-align: left;
|
||||
margin:0 15px 0 0px;
|
||||
font-weight: 400;
|
||||
color:var(--descrip-text);
|
||||
}
|
||||
.easyuse-model-info p {
|
||||
margin: 5px 0;
|
||||
}
|
||||
.easyuse-model-info a {
|
||||
color: var(--theme-color-light);
|
||||
}
|
||||
.easyuse-model-info a:hover {
|
||||
text-decoration: underline;
|
||||
}
|
||||
.easyuse-model-tags-list {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
list-style: none;
|
||||
gap: 10px;
|
||||
max-height: 200px;
|
||||
overflow: auto;
|
||||
margin: 10px 0;
|
||||
padding: 0;
|
||||
}
|
||||
.easyuse-model-tag {
|
||||
background-color: var(--comfy-input-bg);
|
||||
border: 2px solid var(--border-color);
|
||||
color: var(--input-text);
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 5px;
|
||||
border-radius: 5px;
|
||||
padding: 2px 5px;
|
||||
cursor: pointer;
|
||||
}
|
||||
.easyuse-model-tag--selected span::before {
|
||||
content: "✅";
|
||||
position: absolute;
|
||||
background-color: var(--theme-color-light);
|
||||
left: 0;
|
||||
top: 0;
|
||||
right: 0;
|
||||
bottom: 0;
|
||||
text-align: center;
|
||||
}
|
||||
.easyuse-model-tag:hover {
|
||||
border: 2px solid var(--theme-color-light);
|
||||
}
|
||||
.easyuse-model-tag p {
|
||||
margin: 0;
|
||||
}
|
||||
.easyuse-model-tag span {
|
||||
text-align: center;
|
||||
border-radius: 5px;
|
||||
background-color: var(--theme-color-light);
|
||||
padding: 2px;
|
||||
position: relative;
|
||||
min-width: 20px;
|
||||
overflow: hidden;
|
||||
color: #fff;
|
||||
}
|
||||
|
||||
.easyuse-model-metadata .comfy-modal-content {
|
||||
max-width: 100%;
|
||||
}
|
||||
.easyuse-model-metadata label {
|
||||
margin-right: 1ch;
|
||||
color: #ccc;
|
||||
}
|
||||
|
||||
.easyuse-model-metadata span {
|
||||
color: var(--theme-color-light);
|
||||
}
|
||||
|
||||
.easyuse-preview {
|
||||
max-width:660px;
|
||||
margin-right: 15px;
|
||||
position: relative;
|
||||
}
|
||||
.easyuse-preview-group{
|
||||
position: relative;
|
||||
overflow: hidden;
|
||||
border-radius:.5rem;
|
||||
width: 660px;
|
||||
}
|
||||
.easyuse-preview-list{
|
||||
display: flex;
|
||||
flex-wrap: nowrap;
|
||||
width: 100%;
|
||||
transition: all .5s ease-in-out;
|
||||
}
|
||||
.easyuse-preview-list.no-transition{
|
||||
transition: none;
|
||||
}
|
||||
.easyuse-preview-slide{
|
||||
display: flex;
|
||||
flex-basis: calc(50% - 5px);
|
||||
flex-grow: 0;
|
||||
flex-shrink: 0;
|
||||
position: relative;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
padding-right:5px;
|
||||
padding-left:0;
|
||||
}
|
||||
.easyuse-preview-slide:nth-child(even){
|
||||
padding-left:5px;
|
||||
padding-right:0;
|
||||
}
|
||||
.easyuse-preview-slide-content{
|
||||
position: relative;
|
||||
min-height:150px;
|
||||
width: 100%;
|
||||
}
|
||||
.easyuse-preview-slide-content .save{
|
||||
position: absolute;
|
||||
right: 6px;
|
||||
z-index: 12;
|
||||
bottom: 6px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
height: 26px;
|
||||
padding: 0 9px;
|
||||
color: var(--input-text);
|
||||
font-size: 12px;
|
||||
line-height: 26px;
|
||||
background: rgba(0, 0, 0, .5);
|
||||
border-radius: 13px;
|
||||
cursor: pointer;
|
||||
min-width:80px;
|
||||
text-align: center;
|
||||
}
|
||||
.easyuse-preview-slide-content .save:hover{
|
||||
filter: brightness(120%);
|
||||
will-change: auto;
|
||||
}
|
||||
|
||||
.easyuse-preview-slide-content img {
|
||||
border-radius: 14px;
|
||||
object-position: center center;
|
||||
max-width: 100%;
|
||||
max-height:700px;
|
||||
border-style: none;
|
||||
vertical-align: middle;
|
||||
}
|
||||
.easyuse-preview button {
|
||||
position: absolute;
|
||||
z-index:10;
|
||||
top: 50%;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
width:30px;
|
||||
height:30px;
|
||||
border-radius:15px;
|
||||
border:1px solid rgba(66, 63, 78, .15);
|
||||
background-color: rgba(66, 63, 78, .5);
|
||||
color:hsla(0, 0%, 100%, .8);
|
||||
transition-property: color, background-color, border-color, text-decoration-color, fill, stroke;
|
||||
transition-timing-function: cubic-bezier(.4,0,.2,1);
|
||||
transition-duration: .15s;
|
||||
transform: translateY(-50%);
|
||||
}
|
||||
.easyuse-preview button.left{
|
||||
left:10px;
|
||||
}
|
||||
.easyuse-preview button.right{
|
||||
right:10px;
|
||||
}
|
||||
|
||||
.easyuse-model-detail{
|
||||
margin-top: 16px;
|
||||
overflow: hidden;
|
||||
border: 1px solid var(--border-color);
|
||||
border-radius: 8px;
|
||||
width:300px;
|
||||
}
|
||||
.easyuse-model-detail-head{
|
||||
height: 40px;
|
||||
padding: 0 10px;
|
||||
font-weight: 500;
|
||||
font-size: 14px;
|
||||
font-style: normal;
|
||||
line-height: 40px;
|
||||
}
|
||||
.easyuse-model-detail-body{
|
||||
box-sizing: border-box;
|
||||
font-size: 12px;
|
||||
}
|
||||
.easyuse-model-detail-item{
|
||||
display: flex;
|
||||
justify-content: flex-start;
|
||||
border-top: 1px solid var(--border-color);
|
||||
}
|
||||
.easyuse-model-detail-item-label{
|
||||
flex-shrink: 0;
|
||||
width: 88px;
|
||||
padding-top: 5px;
|
||||
padding-bottom: 5px;
|
||||
padding-left: 10px;
|
||||
border-right: 1px solid var(--border-color);
|
||||
color: var(--input-text);
|
||||
font-weight: 400;
|
||||
}
|
||||
.easyuse-model-detail-item-value{
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
padding: 5px 10px 5px 10px;
|
||||
color: var(--input-text);
|
||||
}
|
||||
.easyuse-model-detail-textarea{
|
||||
border-top:1px solid var(--border-color);
|
||||
padding:10px;
|
||||
height:100px;
|
||||
overflow-y: auto;
|
||||
font-size: 12px;
|
||||
}
|
||||
.easyuse-model-detail-textarea textarea{
|
||||
width:100%;
|
||||
height:100%;
|
||||
border:0;
|
||||
background-color:transparent;
|
||||
color: var(--input-text);
|
||||
}
|
||||
.easyuse-model-detail-textarea textarea::placeholder{
|
||||
color:var(--descrip-text);
|
||||
}
|
||||
.easyuse-model-detail-textarea.empty{
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
color: var(--descrip-text);
|
||||
}
|
||||
|
||||
.easyuse-model-notes {
|
||||
background-color: rgba(0, 0, 0, 0.25);
|
||||
padding: 5px;
|
||||
margin-top: 5px;
|
||||
}
|
||||
.easyuse-model-notes:empty {
|
||||
display: none;
|
||||
}
|
||||
@@ -1,3 +1,6 @@
|
||||
.easyuse-prompt-styles{
|
||||
overflow: auto;
|
||||
}
|
||||
.easyuse-prompt-styles .tools{
|
||||
display:flex;
|
||||
justify-content:space-between;
|
||||
@@ -41,9 +44,13 @@
|
||||
min-height: 150px;
|
||||
height: calc(100% - 40px);
|
||||
overflow: auto;
|
||||
// display: flex;
|
||||
// flex-wrap: wrap;
|
||||
/*display: flex;*/
|
||||
/*flex-wrap: wrap;*/
|
||||
}
|
||||
.easyuse-prompt-styles-list.no-top{
|
||||
height: auto;
|
||||
}
|
||||
|
||||
.easyuse-prompt-styles-tag{
|
||||
display: inline-block;
|
||||
vertical-align: middle;
|
||||
|
||||
+4
-1
@@ -1,5 +1,8 @@
|
||||
:root {
|
||||
--theme-color:#3f3eed;
|
||||
--theme-color-light:#006691;
|
||||
--theme-color-light: #008ecb;
|
||||
--success-color: #52c41a;
|
||||
--error-color: #ff4d4f;
|
||||
--warning-color: #faad14;
|
||||
--font-family: Inter, -apple-system, BlinkMacSystemFont, Helvetica Neue, sans-serif;
|
||||
}
|
||||
@@ -0,0 +1,110 @@
|
||||
.easyuse-toast-container{
|
||||
position: fixed;
|
||||
z-index: 99999;
|
||||
top: 0;
|
||||
left: 0;
|
||||
width: 100%;
|
||||
height: 0;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
justify-content: start;
|
||||
padding:10px 0;
|
||||
}
|
||||
.easyuse-toast-container > div {
|
||||
position: relative;
|
||||
height: fit-content;
|
||||
padding: 4px;
|
||||
margin-top: -100px; /* re-set by JS */
|
||||
opacity: 0;
|
||||
transition: all 0.33s ease-in-out;
|
||||
z-index: 3;
|
||||
}
|
||||
|
||||
.easyuse-toast-container > div:last-child {
|
||||
z-index: 2;
|
||||
}
|
||||
|
||||
.easyuse-toast-container > div:not(.-show) {
|
||||
z-index: 1;
|
||||
}
|
||||
|
||||
.easyuse-toast-container > div.-show {
|
||||
opacity: 1;
|
||||
margin-top: 0px !important;
|
||||
}
|
||||
|
||||
.easyuse-toast-container > div.-show {
|
||||
opacity: 1;
|
||||
transform: translateY(0%);
|
||||
}
|
||||
|
||||
.easyuse-toast-container > div > div {
|
||||
position: relative;
|
||||
background: var(--comfy-menu-bg);
|
||||
color: var(--input-text);
|
||||
display: flex;
|
||||
flex-direction: row;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
height: fit-content;
|
||||
box-shadow: 0 0 10px rgba(0, 0, 0, 0.88);
|
||||
padding: 9px 12px;
|
||||
border-radius: 8px;
|
||||
font-family: Arial, sans-serif;
|
||||
font-size: 14px;
|
||||
pointer-events: all;
|
||||
}
|
||||
|
||||
.easyuse-toast-container > div > div > span {
|
||||
display: flex;
|
||||
flex-direction: row;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
.easyuse-toast-container > div > div > span svg {
|
||||
width: 16px;
|
||||
height: auto;
|
||||
margin-right: 8px;
|
||||
}
|
||||
|
||||
.easyuse-toast-container > div > div > span svg[data-icon=info-circle]{
|
||||
fill: var(--theme-color-light);
|
||||
}
|
||||
.easyuse-toast-container > div > div > span svg[data-icon=check-circle]{
|
||||
fill: var(--success-color);
|
||||
}
|
||||
.easyuse-toast-container > div > div > span svg[data-icon=close-circle]{
|
||||
fill: var(--error-color);
|
||||
}
|
||||
.easyuse-toast-container > div > div > span svg[data-icon=exclamation-circle]{
|
||||
fill: var(--warning-color);
|
||||
}
|
||||
/*rotate animation*/
|
||||
@keyframes rotate {
|
||||
0% {
|
||||
transform: rotate(0deg);
|
||||
}
|
||||
100% {
|
||||
transform: rotate(360deg);
|
||||
}
|
||||
}
|
||||
.easyuse-toast-container > div > div > span svg[data-icon=loading]{
|
||||
fill: var(--theme-color);
|
||||
animation: rotate 1s linear infinite;
|
||||
}
|
||||
|
||||
.easyuse-toast-container a {
|
||||
cursor: pointer;
|
||||
text-decoration: underline;
|
||||
color: var(--theme-color-light);
|
||||
margin-left: 4px;
|
||||
display: inline-block;
|
||||
line-height: 1;
|
||||
}
|
||||
|
||||
.easyuse-toast-container a:hover {
|
||||
color: var(--theme-color-light);
|
||||
text-decoration: none;
|
||||
}
|
||||
@@ -0,0 +1,213 @@
|
||||
.easyuse-toolbar{
|
||||
background: rgba(15,15,15,.5);
|
||||
backdrop-filter: blur(4px) brightness(120%);
|
||||
border-radius:0 12px 12px 0;
|
||||
min-width:50px;
|
||||
height:24px;
|
||||
position: fixed;
|
||||
bottom:85px;
|
||||
left:0px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
z-index:10000;
|
||||
}
|
||||
.easyuse-toolbar.disable-render-info{
|
||||
bottom: 55px;
|
||||
}
|
||||
.easyuse-toolbar-item{
|
||||
border-radius:20px;
|
||||
height: 20px;
|
||||
width:20px;
|
||||
cursor: pointer;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
transition: all 0.3s ease-in-out;
|
||||
margin-left:2.5px;
|
||||
}
|
||||
.easyuse-toolbar-icon{
|
||||
width: 14px;
|
||||
height: 14px;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
font-size: 12px;
|
||||
color:white;
|
||||
transition: all 0.3s ease-in-out;
|
||||
}
|
||||
.easyuse-toolbar-tips{
|
||||
visibility: hidden;
|
||||
opacity: 0;
|
||||
position: absolute;
|
||||
top: -25px;
|
||||
left: 0;
|
||||
color: var(--descrip-text);
|
||||
padding: 2px 5px;
|
||||
border-radius: 5px;
|
||||
font-size: 11px;
|
||||
min-width:100px;
|
||||
transition: all 0.3s ease-in-out;
|
||||
}
|
||||
.easyuse-toolbar-item:hover{
|
||||
background:rgba(12,12,12,1);
|
||||
}
|
||||
.easyuse-toolbar-item:hover .easyuse-toolbar-tips{
|
||||
opacity: 1;
|
||||
visibility: visible;
|
||||
}
|
||||
.easyuse-toolbar-item:hover .easyuse-toolbar-icon.group{
|
||||
color:var(--warning-color);
|
||||
}
|
||||
.easyuse-toolbar-item:hover .easyuse-toolbar-icon.rocket{
|
||||
color:var(--theme-color-light);
|
||||
}
|
||||
.easyuse-toolbar-item:hover .easyuse-toolbar-icon.question{
|
||||
color:var(--success-color);
|
||||
}
|
||||
|
||||
|
||||
.easyuse-guide-dialog{
|
||||
max-width: 300px;
|
||||
font-family: var(--font-family);
|
||||
position: absolute;
|
||||
z-index:100;
|
||||
left:0;
|
||||
bottom:140px;
|
||||
background: rgba(25,25,25,.85);
|
||||
backdrop-filter: blur(8px) brightness(120%);
|
||||
border-radius:0 12px 12px 0;
|
||||
padding:10px;
|
||||
transition: .5s all ease-in-out;
|
||||
visibility: visible;
|
||||
opacity: 1;
|
||||
transform: translateX(0%);
|
||||
}
|
||||
.easyuse-guide-dialog.disable-render-info{
|
||||
bottom:110px;
|
||||
}
|
||||
.easyuse-guide-dialog-top{
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
}
|
||||
.easyuse-guide-dialog-top .icon{
|
||||
width: 12px;
|
||||
height:12px;
|
||||
}
|
||||
.easyuse-guide-dialog.hidden{
|
||||
opacity: 0;
|
||||
transform: translateX(-50%);
|
||||
visibility: hidden;
|
||||
}
|
||||
.easyuse-guide-dialog .closeBtn{
|
||||
float: right;
|
||||
color: var(--input-text);
|
||||
border-radius:30px;
|
||||
background-color: var(--comfy-input-bg);
|
||||
border: 1px solid var(--border-color);
|
||||
cursor: pointer;
|
||||
aspect-ratio: 1 / 1;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
}
|
||||
.easyuse-guide-dialog .closeBtn:hover{
|
||||
filter:brightness(120%);
|
||||
}
|
||||
.easyuse-guide-dialog-title{
|
||||
color:var(--input-text);
|
||||
font-size: 16px;
|
||||
font-weight: bold;
|
||||
margin-bottom: 5px;
|
||||
}
|
||||
.easyuse-guide-dialog-remark{
|
||||
color: var(--input-text);
|
||||
font-size: 12px;
|
||||
margin-top: 5px;
|
||||
}
|
||||
.easyuse-guide-dialog-content{
|
||||
max-height: 600px;
|
||||
overflow: auto;
|
||||
}
|
||||
.easyuse-guide-dialog a, .easyuse-guide-dialog a:visited{
|
||||
color: var(--theme-color-light);
|
||||
cursor: pointer;
|
||||
}
|
||||
.easyuse-guide-dialog-note{
|
||||
margin-top: 20px;
|
||||
color:white;
|
||||
}
|
||||
.easyuse-guide-dialog p{
|
||||
margin:4px 0;
|
||||
font-size: 12px;
|
||||
font-weight: 300;
|
||||
}
|
||||
.markdown-body h1, .markdown-body h2, .markdown-body h3, .markdown-body h4, .markdown-body h5, .markdown-body h6 {
|
||||
margin-top: 12px;
|
||||
margin-bottom: 8px;
|
||||
font-weight: 600;
|
||||
line-height: 1.25;
|
||||
padding-bottom: 5px;
|
||||
border-bottom: 1px solid var(--border-color);
|
||||
color: var(--input-text);
|
||||
}
|
||||
.markdown-body h1{
|
||||
font-size: 18px;
|
||||
}
|
||||
.markdown-body h2{
|
||||
font-size: 16px;
|
||||
}
|
||||
.markdown-body h3{
|
||||
font-size: 14px;
|
||||
}
|
||||
.markdown-body h4{
|
||||
font-size: 13px;
|
||||
}
|
||||
.markdown-body table {
|
||||
display: block;
|
||||
/*width: 100%;*/
|
||||
/*width: max-content;*/
|
||||
max-width: 300px;
|
||||
overflow: auto;
|
||||
color:var(--input-text);
|
||||
box-sizing: border-box;
|
||||
border: 1px solid var(--border-color);
|
||||
text-align: left;
|
||||
width: 100%;
|
||||
}
|
||||
.markdown-body table th, .markdown-body table td {
|
||||
padding: 6px 13px;
|
||||
font-size: 12px;
|
||||
margin:0;
|
||||
border-right: 1px solid var(--border-color);
|
||||
border-bottom: 1px solid var(--border-color);
|
||||
}
|
||||
.markdown-body table td {
|
||||
font-size: 12px;
|
||||
}
|
||||
.markdown-body table th:last-child, .markdown-body table td:last-child{
|
||||
border-right: none;
|
||||
}
|
||||
.markdown-body table tr:last-child td{
|
||||
border-bottom: none;
|
||||
}
|
||||
.markdown-body table th{
|
||||
font-weight: bold;
|
||||
width: auto;
|
||||
min-width: 70px;
|
||||
}
|
||||
.markdown-body table th:last-child{
|
||||
width:100%;
|
||||
}
|
||||
.markdown-body .warning{
|
||||
color:var(--warning-color)
|
||||
}
|
||||
.markdown-body .error{
|
||||
color:var(--error-color)
|
||||
}
|
||||
.markdown-body .success{
|
||||
color:var(--success-color)
|
||||
}
|
||||
.markdown-body .link{
|
||||
color:var(--theme-color-light)
|
||||
}
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
import { app } from "/scripts/app.js";
|
||||
import { app } from "../../../scripts/app.js";
|
||||
|
||||
|
||||
app.registerExtension({
|
||||
|
||||
+87
-5
@@ -2,18 +2,100 @@ import {getLocale} from './utils.js'
|
||||
const locale = getLocale()
|
||||
|
||||
const zhCN = {
|
||||
"Workflow created by": "工作流创建者",
|
||||
"Watch more video content": "观看更多视频内容",
|
||||
"Workflow Guide":"工作流指南",
|
||||
// ExtraMenu
|
||||
"💎 View Checkpoint Info...": "💎 查看 Checkpoint 信息...",
|
||||
"💎 View Lora Info...": "💎 查看 Lora 信息...",
|
||||
"🔃 Reload Node": "🔃 刷新节点",
|
||||
// ModelInfo
|
||||
"Updated At:": "最近更新:",
|
||||
"Created At:": "首次发布:",
|
||||
"✏️ Edit": "✏️ 编辑",
|
||||
"💾 Save": "💾 保存",
|
||||
"No notes": "当前还没有备注内容",
|
||||
"Saving Notes...": "正在保存备注...",
|
||||
"Type your notes here":"在这里输入备注内容",
|
||||
"ModelName":"模型名称",
|
||||
"Models Required":"所需模型",
|
||||
"Download Model": "下载模型",
|
||||
"Source Url": "模型源地址",
|
||||
"Notes": "备注",
|
||||
"Type": "类型",
|
||||
"Trained Words": "训练词",
|
||||
"BaseModel": "基础算法",
|
||||
"Details": "详情",
|
||||
"Description": "描述",
|
||||
"Download": "下载量",
|
||||
"Source": "来源",
|
||||
"Saving Preview...": "正在保存预览图...",
|
||||
"Saving Succeed":"保存成功",
|
||||
"Clean SuccessFully":"清理成功",
|
||||
"Clean Failed": "清理失败",
|
||||
"Saving Failed":"保存失败",
|
||||
"No COMBO link": "沒有找到COMBO连接",
|
||||
"Reboot ComfyUI":"重启ComfyUI",
|
||||
"Are you sure you'd like to reboot the server?": "是否要重启ComfyUI?",
|
||||
// GroupMap
|
||||
"Groups Map (EasyUse)": "管理组 (EasyUse)",
|
||||
"Groups Map": "管理组",
|
||||
"Cleanup Of GPU Usage": "清理GPU占用",
|
||||
"Please stop all running tasks before cleaning GPU": "请在清理GPU之前停止所有运行中的任务",
|
||||
"Always": "启用中",
|
||||
"Bypass": "已忽略",
|
||||
"Never": "已停用",
|
||||
"Auto Sorting": "自动排序",
|
||||
"Toggle `Show/Hide` can set mode of group, LongPress can set group nodes to never": "点击`启用中/已忽略`可设置组模式, 长按可停用该组节点",
|
||||
// Quick
|
||||
"Enable ALT+1~9 to paste nodes from nodes template (ComfyUI-Easy-Use)": "启用ALT1~9从节点模板粘贴到工作流(ComfyUI-Easy-Use)",
|
||||
"Enable process bar in queue button (ComfyUI-Easy-Use)": "启用提示词队列进度显示条(ComfyUI-Easy-Use)"
|
||||
"Enable ALT+1~9 to paste nodes from nodes template (ComfyUI-Easy-Use)": "启用ALT1~9从节点模板粘贴到工作流 (ComfyUI-Easy-Use)",
|
||||
"Enable process bar in queue button (ComfyUI-Easy-Use)": "启用提示词队列进度显示条 (ComfyUI-Easy-Use)",
|
||||
"Enable ContextMenu Auto Nest Subdirectories (ComfyUI-Easy-Use)": "启用上下文菜单自动嵌套子目录 (ComfyUI-Easy-Use)",
|
||||
"Enable tool bar fixed on the left-bottom (ComfyUI-Easy-Use)": "启用工具栏固定在左下角 (ComfyUI-Easy-Use)",
|
||||
"Too many thumbnails, have closed the display": "模型缩略图太多啦,为您关闭了显示",
|
||||
// selector
|
||||
"Empty All": "清空所有",
|
||||
"🔎 Type here to search styles ...": "🔎 在此处输入以搜索样式 ...",
|
||||
// account
|
||||
"Loading UserInfo...": "正在获取用户信息...",
|
||||
"Please set the APIKEY first": "请先设置APIKEY",
|
||||
"Setting APIKEY": "设置APIKEY",
|
||||
"Save Account Info": "保存账号信息",
|
||||
"Choose": "选择",
|
||||
"Delete": "删除",
|
||||
"Edit": "编辑",
|
||||
"At least one account is required": "删除失败: 至少需要一个账户",
|
||||
"APIKEY is not Empty": "APIKEY 不能为空",
|
||||
"Add Account": "添加账号",
|
||||
"Getting Your APIKEY": "获取您的APIKEY",
|
||||
// choosers
|
||||
"Choose Selected Images": "选择选中的图片",
|
||||
"Choose images to continue": "选择图片以继续",
|
||||
// seg
|
||||
"Background": "背景",
|
||||
"Hat": "帽子",
|
||||
"Hair": "头发",
|
||||
"Body": "身体",
|
||||
"Face": "脸部",
|
||||
"Clothes": "衣服",
|
||||
"Others": "其他",
|
||||
"Glove": "手套",
|
||||
"Sunglasses": "太阳镜",
|
||||
"Upper-clothes": "上衣",
|
||||
"Dress": "连衣裙",
|
||||
"Coat": "外套",
|
||||
"Socks": "袜子",
|
||||
"Pants": "裤子",
|
||||
"Jumpsuits": "连体衣",
|
||||
"Scarf": "围巾",
|
||||
"Skirt": "裙子",
|
||||
"Left-arm": "左臂",
|
||||
"Right-arm": "右臂",
|
||||
"Left-leg": "左腿",
|
||||
"Right-leg": "右腿",
|
||||
"Left-shoe": "左鞋",
|
||||
"Right-shoe": "右鞋",
|
||||
}
|
||||
|
||||
export const $t = (key) => {
|
||||
return locale === 'zh-CN' ? zhCN[key] : key
|
||||
const cn = zhCN[key]
|
||||
return locale === 'zh-CN' && cn ? cn : key
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
export const quesitonIcon = `<svg t="1714564780771" class="icon" viewBox="0 0 1024 1024" version="1.1" xmlns="http://www.w3.org/2000/svg" p-id="1489" width="200" height="200" data-spm-anchor-id="a313x.search_index.0.i2.5a663a81pw6qup"><path d="M514.048 54.272q95.232 0 178.688 36.352t145.92 98.304 98.304 145.408 35.84 178.688-35.84 178.176-98.304 145.408-145.92 98.304-178.688 35.84-178.176-35.84-145.408-98.304-98.304-145.408-35.84-178.176 35.84-178.688 98.304-145.408 145.408-98.304 178.176-36.352zM515.072 826.368q26.624 0 44.544-17.92t17.92-43.52q0-26.624-17.92-44.544t-44.544-17.92-44.544 17.92-17.92 44.544q0 25.6 17.92 43.52t44.544 17.92zM567.296 574.464q-1.024-16.384 20.48-34.816t48.128-40.96 49.152-50.688 24.576-65.024q2.048-39.936-8.192-74.752t-33.792-59.904-60.928-39.936-87.552-14.848q-62.464 0-103.936 22.016t-67.072 53.248-35.84 64.512-9.216 55.808q1.024 26.624 16.896 38.912t34.304 12.8 33.792-10.24 15.36-31.232q0-12.288 7.68-30.208t20.992-34.304 32.256-27.648 42.496-11.264q46.08 0 73.728 23.04t25.6 57.856q0 17.408-10.24 32.256t-26.112 28.672-33.792 27.648-33.792 28.672-26.624 32.256-11.776 37.888l1.024 38.912q0 15.36 14.336 29.184t37.888 14.848q23.552-1.024 37.376-15.36t12.8-32.768l0-24.576z" p-id="1490" fill="currentColor"></path></svg>`
|
||||
export const rocketIcon = `<svg t="1714565020764" class="icon" viewBox="0 0 1024 1024" version="1.1" xmlns="http://www.w3.org/2000/svg" p-id="7999" width="200" height="200"><path d="M810.438503 379.664884l-71.187166-12.777183C737.426025 180.705882 542.117647 14.602496 532.991087 7.301248c-12.777184-10.951872-32.855615-10.951872-47.45811 0-9.12656 7.301248-204.434938 175.229947-206.26025 359.586453l-67.536542 10.951871c-18.253119 3.650624-31.030303 18.253119-31.030303 36.506239v189.832442c0 10.951872 5.475936 21.903743 12.777184 27.379679 7.301248 5.475936 14.602496 9.12656 23.729055 9.12656h5.475936l133.247772-23.729055c40.156863 47.458111 91.265597 73.012478 151.500891 73.012477 60.235294 0 111.344029-27.379679 151.500891-74.837789l136.898396 23.729055h5.475936c9.12656 0 16.427807-3.650624 23.729055-9.12656 9.12656-7.301248 12.777184-16.427807 12.777184-27.379679V412.520499c1.825312-14.602496-10.951872-29.204991-27.379679-32.855615zM620.606061 766.631016H401.568627c-20.078431 0-36.506239 16.427807-36.506238 36.506239v109.518716c0 14.602496 9.12656 29.204991 23.729055 34.680927 14.602496 5.475936 31.030303 1.825312 40.156863-9.126559l16.427807-18.25312 32.855615 80.313726c5.475936 14.602496 18.253119 23.729055 34.680927 23.729055 16.427807 0 27.379679-9.12656 34.680927-23.729055l32.855615-80.313726 16.427807 18.25312c10.951872 10.951872 25.554367 14.602496 40.156863 9.126559 14.602496-5.475936 23.729055-18.253119 23.729055-34.680927v-109.518716c-3.650624-20.078431-20.078431-36.506239-40.156862-36.506239z" fill="currentColor" p-id="8000"></path></svg>`
|
||||
export const groupIcon = `<svg t="1714565543756" class="icon" viewBox="0 0 1024 1024" version="1.1" xmlns="http://www.w3.org/2000/svg" p-id="22538" width="200" height="200"><path d="M871.616 64H152.384c-31.488 0-60.416 25.28-60.416 58.24v779.52c0 32.896 26.24 58.24 60.352 58.24h719.232c34.112 0 60.352-25.344 60.352-58.24V122.24c0.128-32.96-28.8-58.24-60.288-58.24zM286.272 512c-23.616 0-44.672-20.224-44.672-43.008 0-22.784 20.992-43.008 44.608-43.008 23.616 0 44.608 20.224 44.608 43.008A43.328 43.328 0 0 1 286.272 512z m0-202.496c-23.616 0-44.608-20.224-44.608-43.008 0-22.784 20.992-43.008 44.608-43.008 23.616 0 44.608 20.224 44.608 43.008a43.456 43.456 0 0 1-44.608 43.008zM737.728 512H435.904c-23.68 0-44.672-20.224-44.672-43.008 0-22.784 20.992-43.008 44.608-43.008h299.264c23.616 0 44.608 20.224 44.608 43.008a42.752 42.752 0 0 1-41.984 43.008z m0-202.496H435.904c-23.616 0-44.608-20.224-44.608-43.008 0-22.784 20.992-43.008 44.608-43.008h299.264c23.616 0 44.608 20.224 44.608 43.008a42.88 42.88 0 0 1-42.048 43.008z" p-id="22539" fill="currentColor"></path></svg>`
|
||||
export const rebootIcon = `<svg t="1714568501931" class="icon" viewBox="0 0 1024 1024" version="1.1" xmlns="http://www.w3.org/2000/svg" p-id="4275" width="200" height="200"><path d="M511.721751 0.000278a511.999861 511.999861 0 1 0 512.277971 511.721751A511.721751 511.721751 0 0 0 511.721751 0.000278zM184.386696 511.722029A36.988583 36.988583 0 0 1 222.487718 475.011556h92.888622a36.710473 36.710473 0 0 1 0 73.420947H222.487718a36.710473 36.710473 0 0 1-38.101022-36.710474z m201.351385 158.522499l-65.911986 65.911987a36.988583 36.988583 0 0 1-62.852781-25.864197 38.101022 38.101022 0 0 1 10.846276-26.142307L333.731577 618.238024a36.710473 36.710473 0 1 1 52.006504 52.006504z m29.201513-256.138985a36.710473 36.710473 0 0 1-52.006504 0l-65.633877-65.633877a36.988583 36.988583 0 0 1 26.142307-62.85278 36.154254 36.154254 0 0 1 25.864197 10.846276L414.939594 361.54282a36.988583 36.988583 0 0 1 0 52.562723z m135.439398 373.779366a37.266693 37.266693 0 0 1-36.988583 36.988583 36.988583 36.988583 0 0 1-36.710473-36.988583V695.274397a36.988583 36.988583 0 0 1 36.710473-36.988583A37.266693 37.266693 0 0 1 550.378992 695.274397z m0-459.437137a37.266693 37.266693 0 0 1-36.988583 36.988583 36.988583 36.988583 0 0 1-36.710473-36.988583V235.559149a36.988583 36.988583 0 0 1 36.710473-36.988583 37.544802 37.544802 0 0 1 36.988583 36.988583z m63.965219 15.85225L679.978088 278.109926a36.710473 36.710473 0 0 1 52.006504 51.728394L667.463154 396.584635a37.544802 37.544802 0 0 1-52.284614 0 36.988583 36.988583 0 0 1-10.568166-26.142306 36.432364 36.432364 0 0 1 9.733837-26.142307z m122.090135 397.974905a37.544802 37.544802 0 0 1-52.284613 0l-65.355767-65.911986a36.154254 36.154254 0 0 1 0-51.728395 36.710473 36.710473 0 0 1 25.864197-10.846276 35.876145 35.876145 0 0 1 25.864197 10.846276l65.911986 65.633877a36.988583 36.988583 0 0 1 0 52.006504z m66.468206-194.676753h-92.888622a36.710473 36.710473 0 0 1 0-73.420947h92.888622a36.710473 36.710473 0 0 1 0 73.420947z" fill="currentColor" p-id="4276"></path></svg>`
|
||||
export const closeIcon = `<svg t="1714965640187" class="icon" viewBox="0 0 1024 1024" version="1.1" xmlns="http://www.w3.org/2000/svg" p-id="4264" width="200" height="200"><path d="M597.795527 511.488347 813.564755 295.718095c23.833825-23.833825 23.833825-62.47489 0.001023-86.307691-23.832801-23.832801-62.47489-23.833825-86.307691 0L511.487835 425.180656 295.717583 209.410404c-23.833825-23.833825-62.475913-23.833825-86.307691 0-23.832801 23.832801-23.833825 62.47489 0 86.308715l215.769228 215.769228L209.410915 727.258599c-23.833825 23.833825-23.833825 62.47489 0 86.307691 23.832801 23.833825 62.473867 23.833825 86.307691 0l215.768205-215.768205 215.769228 215.769228c23.834848 23.833825 62.475913 23.832801 86.308715 0 23.833825-23.833825 23.833825-62.47489 0-86.307691L597.795527 511.488347z" fill="currentColor" p-id="4265"></path></svg>`
|
||||
@@ -0,0 +1,683 @@
|
||||
import { $el, ComfyDialog } from "../../../../scripts/ui.js";
|
||||
import { api } from "../../../../scripts/api.js";
|
||||
import {formatTime} from './utils.js';
|
||||
import {$t} from "./i18n.js";
|
||||
import {toast} from "./toast.js";
|
||||
|
||||
class MetadataDialog extends ComfyDialog {
|
||||
constructor() {
|
||||
super();
|
||||
this.element.classList.add("easyuse-model-metadata");
|
||||
}
|
||||
show(metadata) {
|
||||
super.show(
|
||||
$el(
|
||||
"div",
|
||||
Object.keys(metadata).map((k) =>
|
||||
$el("div", [$el("label", { textContent: k }), $el("span", { textContent: metadata[k] })])
|
||||
)
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
export class ModelInfoDialog extends ComfyDialog {
|
||||
constructor(name) {
|
||||
super();
|
||||
this.name = name;
|
||||
this.element.classList.add("easyuse-model-info");
|
||||
}
|
||||
|
||||
get customNotes() {
|
||||
return this.metadata["easyuse.notes"];
|
||||
}
|
||||
|
||||
set customNotes(v) {
|
||||
this.metadata["easyuse.notes"] = v;
|
||||
}
|
||||
|
||||
get hash() {
|
||||
return this.metadata["easyuse.sha256"];
|
||||
}
|
||||
|
||||
async show(type, value) {
|
||||
this.type = type;
|
||||
|
||||
const req = api.fetchApi("/easyuse/metadata/" + encodeURIComponent(`${type}/${value}`));
|
||||
this.info = $el("div", { style: { flex: "auto" } });
|
||||
// this.img = $el("img", { style: { display: "none" } });
|
||||
this.imgCurrent = 0
|
||||
this.imgList = $el("div.easyuse-preview-list",{
|
||||
style: { display: "none" }
|
||||
})
|
||||
this.imgWrapper = $el("div.easyuse-preview", [
|
||||
$el("div.easyuse-preview-group",[
|
||||
this.imgList
|
||||
]),
|
||||
]);
|
||||
this.main = $el("main", { style: { display: "flex" } }, [this.imgWrapper, this.info]);
|
||||
this.content = $el("div.easyuse-model-content", [
|
||||
$el("div.easyuse-model-header",[$el("h2", { textContent: this.name })])
|
||||
, this.main]);
|
||||
|
||||
const loading = $el("div", { textContent: "ℹ️ Loading...", parent: this.content });
|
||||
|
||||
super.show(this.content);
|
||||
|
||||
this.metadata = await (await req).json();
|
||||
this.viewMetadata.style.cursor = this.viewMetadata.style.opacity = "";
|
||||
this.viewMetadata.removeAttribute("disabled");
|
||||
|
||||
loading.remove();
|
||||
this.addInfo();
|
||||
}
|
||||
|
||||
createButtons() {
|
||||
const btns = super.createButtons();
|
||||
this.viewMetadata = $el("button", {
|
||||
type: "button",
|
||||
textContent: "View raw metadata",
|
||||
disabled: "disabled",
|
||||
style: {
|
||||
opacity: 0.5,
|
||||
cursor: "not-allowed",
|
||||
},
|
||||
onclick: (e) => {
|
||||
if (this.metadata) {
|
||||
new MetadataDialog().show(this.metadata);
|
||||
}
|
||||
},
|
||||
});
|
||||
|
||||
btns.unshift(this.viewMetadata);
|
||||
return btns;
|
||||
}
|
||||
|
||||
parseNote() {
|
||||
if (!this.customNotes) return [];
|
||||
|
||||
let notes = [];
|
||||
// Extract links from notes
|
||||
const r = new RegExp("(\\bhttps?:\\/\\/[^\\s]+)", "g");
|
||||
let end = 0;
|
||||
let m;
|
||||
do {
|
||||
m = r.exec(this.customNotes);
|
||||
let pos;
|
||||
let fin = 0;
|
||||
if (m) {
|
||||
pos = m.index;
|
||||
fin = m.index + m[0].length;
|
||||
} else {
|
||||
pos = this.customNotes.length;
|
||||
}
|
||||
|
||||
let pre = this.customNotes.substring(end, pos);
|
||||
if (pre) {
|
||||
pre = pre.replaceAll("\n", "<br>");
|
||||
notes.push(
|
||||
$el("span", {
|
||||
innerHTML: pre,
|
||||
})
|
||||
);
|
||||
}
|
||||
if (m) {
|
||||
notes.push(
|
||||
$el("a", {
|
||||
href: m[0],
|
||||
textContent: m[0],
|
||||
target: "_blank",
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
end = fin;
|
||||
} while (m);
|
||||
return notes;
|
||||
}
|
||||
|
||||
addInfoEntry(name, value) {
|
||||
return $el(
|
||||
"p",
|
||||
{
|
||||
parent: this.info,
|
||||
},
|
||||
[
|
||||
typeof name === "string" ? $el("label", { textContent: name + ": " }) : name,
|
||||
typeof value === "string" ? $el("span", { textContent: value }) : value,
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
async getCivitaiDetails() {
|
||||
const req = await fetch("https://civitai.com/api/v1/model-versions/by-hash/" + this.hash);
|
||||
if (req.status === 200) {
|
||||
return await req.json();
|
||||
} else if (req.status === 404) {
|
||||
throw new Error("Model not found");
|
||||
} else {
|
||||
throw new Error(`Error loading info (${req.status}) ${req.statusText}`);
|
||||
}
|
||||
}
|
||||
|
||||
addCivitaiInfo() {
|
||||
const promise = this.getCivitaiDetails();
|
||||
const content = $el("span", { textContent: "ℹ️ Loading..." });
|
||||
|
||||
this.addInfoEntry(
|
||||
$el("label", [
|
||||
$el("img", {
|
||||
style: {
|
||||
width: "18px",
|
||||
position: "relative",
|
||||
top: "3px",
|
||||
margin: "0 5px 0 0",
|
||||
},
|
||||
src: "https://civitai.com/favicon.ico",
|
||||
}),
|
||||
$el("span", { textContent: "Civitai: " }),
|
||||
]),
|
||||
content
|
||||
);
|
||||
|
||||
return promise
|
||||
.then((info) => {
|
||||
this.imgWrapper.style.display = 'block'
|
||||
// 变更标题信息
|
||||
let header = this.element.querySelector('.easyuse-model-header')
|
||||
if(header){
|
||||
header.replaceChildren(
|
||||
$el("h2", { textContent: this.name }),
|
||||
$el("div.easyuse-model-header-remark",[
|
||||
$el("h5", { textContent: $t("Updated At:") + formatTime(new Date(info.updatedAt),'yyyy/MM/dd')}),
|
||||
$el("h5", { textContent: $t("Created At:") + formatTime(new Date(info.updatedAt),'yyyy/MM/dd')}),
|
||||
])
|
||||
)
|
||||
}
|
||||
// 替换内容
|
||||
let textarea = null
|
||||
let notes = this.parseNote.call(this)
|
||||
let editText = $t("✏️ Edit")
|
||||
console.log(notes)
|
||||
let textarea_div = $el("div.easyuse-model-detail-textarea",[
|
||||
$el("p",notes?.length>0 ? notes : {textContent:$t('No notes')}),
|
||||
])
|
||||
if(!notes || notes.length == 0) textarea_div.classList.add('empty')
|
||||
else textarea_div.classList.remove('empty')
|
||||
this.info.replaceChildren(
|
||||
$el("div.easyuse-model-detail",[
|
||||
$el("div.easyuse-model-detail-head.flex-b",[
|
||||
$el('span',$t("Notes")),
|
||||
$el("a", {
|
||||
textContent: editText,
|
||||
href: "#",
|
||||
style: {
|
||||
fontSize: "12px",
|
||||
float: "right",
|
||||
color: "var(--warning-color)",
|
||||
textDecoration: "none",
|
||||
},
|
||||
onclick: async (e) => {
|
||||
e.preventDefault();
|
||||
|
||||
if (textarea) {
|
||||
if(textarea.value != this.customNotes){
|
||||
toast.showLoading($t('Saving Notes...'))
|
||||
this.customNotes = textarea.value;
|
||||
const resp = await api.fetchApi(
|
||||
"/easyuse/metadata/notes/" + encodeURIComponent(`${this.type}/${this.name}`),
|
||||
{
|
||||
method: "POST",
|
||||
body: this.customNotes,
|
||||
}
|
||||
);
|
||||
toast.hideLoading()
|
||||
if (resp.status !== 200) {
|
||||
toast.error($t('Saving Failed'))
|
||||
console.error(resp);
|
||||
alert(`Error saving notes (${resp.status}) ${resp.statusText}`);
|
||||
return;
|
||||
}
|
||||
toast.success($t('Saving Succeed'))
|
||||
notes = this.parseNote.call(this)
|
||||
console.log(notes)
|
||||
textarea_div.replaceChildren($el("p",notes?.length>0 ? notes : {textContent:$t('No notes')}));
|
||||
if(textarea.value) textarea_div.classList.remove('empty')
|
||||
else textarea_div.classList.add('empty')
|
||||
}else {
|
||||
textarea_div.replaceChildren($el("p",{textContent:$t('No notes')}));
|
||||
textarea_div.classList.add('empty')
|
||||
}
|
||||
e.target.textContent = editText;
|
||||
textarea.remove();
|
||||
textarea = null;
|
||||
|
||||
} else {
|
||||
e.target.textContent = "💾 Save";
|
||||
textarea = $el("textarea", {
|
||||
placeholder: $t("Type your notes here"),
|
||||
style: {
|
||||
width: "100%",
|
||||
minWidth: "200px",
|
||||
minHeight: "50px",
|
||||
height:"100px"
|
||||
},
|
||||
textContent: this.customNotes,
|
||||
});
|
||||
textarea_div.replaceChildren(textarea);
|
||||
textarea.focus()
|
||||
}
|
||||
}
|
||||
})
|
||||
]),
|
||||
textarea_div
|
||||
]),
|
||||
$el("div.easyuse-model-detail",[
|
||||
$el("div.easyuse-model-detail-head",{textContent:$t("Details")}),
|
||||
$el("div.easyuse-model-detail-body",[
|
||||
$el("div.easyuse-model-detail-item",[
|
||||
$el("div.easyuse-model-detail-item-label",{textContent:$t("Type")}),
|
||||
$el("div.easyuse-model-detail-item-value",{textContent:info.model.type}),
|
||||
]),
|
||||
$el("div.easyuse-model-detail-item",[
|
||||
$el("div.easyuse-model-detail-item-label",{textContent:$t("BaseModel")}),
|
||||
$el("div.easyuse-model-detail-item-value",{textContent:info.baseModel}),
|
||||
]),
|
||||
$el("div.easyuse-model-detail-item",[
|
||||
$el("div.easyuse-model-detail-item-label",{textContent:$t("Download")}),
|
||||
$el("div.easyuse-model-detail-item-value",{textContent:info.stats?.downloadCount || 0}),
|
||||
]),
|
||||
$el("div.easyuse-model-detail-item",[
|
||||
$el("div.easyuse-model-detail-item-label",{textContent:$t("Trained Words")}),
|
||||
$el("div.easyuse-model-detail-item-value",{textContent:info?.trainedWords.join(',') || '-'}),
|
||||
]),
|
||||
$el("div.easyuse-model-detail-item",[
|
||||
$el("div.easyuse-model-detail-item-label",{textContent:$t("Source")}),
|
||||
$el("div.easyuse-model-detail-item-value",[
|
||||
$el("label", [
|
||||
$el("img", {
|
||||
style: {
|
||||
width: "14px",
|
||||
position: "relative",
|
||||
top: "3px",
|
||||
margin: "0 5px 0 0",
|
||||
},
|
||||
src: "https://civitai.com/favicon.ico",
|
||||
}),
|
||||
$el("a", {
|
||||
href: "https://civitai.com/models/" + info.modelId,
|
||||
textContent: "View " + info.model.name,
|
||||
target: "_blank",
|
||||
})
|
||||
])
|
||||
]),
|
||||
])
|
||||
]),
|
||||
])
|
||||
);
|
||||
|
||||
if (info.images?.length) {
|
||||
this.imgCurrent = 0
|
||||
this.isSaving = false
|
||||
info.images.map(cate=>
|
||||
cate.url &&
|
||||
this.imgList.appendChild(
|
||||
$el('div.easyuse-preview-slide',[
|
||||
$el('div.easyuse-preview-slide-content',[
|
||||
$el('img',{src:(cate.url)}),
|
||||
$el("div.save", {
|
||||
textContent: "Save as preview",
|
||||
onclick: async () => {
|
||||
if(this.isSaving) return
|
||||
this.isSaving = true
|
||||
toast.showLoading($t('Saving Preview...'))
|
||||
// Convert the preview to a blob
|
||||
const blob = await (await fetch(cate.url)).blob();
|
||||
|
||||
// Store it in temp
|
||||
const name = "temp_preview." + new URL(cate.url).pathname.split(".")[1];
|
||||
const body = new FormData();
|
||||
body.append("image", new File([blob], name));
|
||||
body.append("overwrite", "true");
|
||||
body.append("type", "temp");
|
||||
|
||||
const resp = await api.fetchApi("/upload/image", {
|
||||
method: "POST",
|
||||
body,
|
||||
});
|
||||
|
||||
if (resp.status !== 200) {
|
||||
this.isSaving = false
|
||||
toast.error($t('Saving Failed'))
|
||||
toast.hideLoading()
|
||||
console.error(resp);
|
||||
alert(`Error saving preview (${req.status}) ${req.statusText}`);
|
||||
return;
|
||||
}
|
||||
|
||||
// Use as preview
|
||||
await api.fetchApi("/easyuse/save/" + encodeURIComponent(`${this.type}/${this.name}`), {
|
||||
method: "POST",
|
||||
body: JSON.stringify({
|
||||
filename: name,
|
||||
type: "temp",
|
||||
}),
|
||||
headers: {
|
||||
"content-type": "application/json",
|
||||
},
|
||||
}).then(_=>{
|
||||
toast.success($t('Saving Succeed'))
|
||||
toast.hideLoading()
|
||||
});
|
||||
this.isSaving = false
|
||||
app.refreshComboInNodes();
|
||||
},
|
||||
})
|
||||
])
|
||||
])
|
||||
)
|
||||
)
|
||||
let _this = this
|
||||
this.imgDistance = (-660 * this.imgCurrent).toString()
|
||||
this.imgList.style.display = ''
|
||||
this.imgList.style.transform = 'translate3d(' + this.imgDistance +'px, 0px, 0px)'
|
||||
this.slides = this.imgList.querySelectorAll('.easyuse-preview-slide')
|
||||
// 添加按钮
|
||||
this.slideLeftButton = $el("button.left",{
|
||||
parent: this.imgWrapper,
|
||||
style:{
|
||||
display:info.images.length <= 2 ? 'none' : 'block'
|
||||
},
|
||||
innerHTML:`<svg viewBox="0 0 15 15" fill="none" xmlns="http://www.w3.org/2000/svg" width="16" height="16" style="transform: rotate(90deg);"><path d="M3.13523 6.15803C3.3241 5.95657 3.64052 5.94637 3.84197 6.13523L7.5 9.56464L11.158 6.13523C11.3595 5.94637 11.6759 5.95657 11.8648 6.15803C12.0536 6.35949 12.0434 6.67591 11.842 6.86477L7.84197 10.6148C7.64964 10.7951 7.35036 10.7951 7.15803 10.6148L3.15803 6.86477C2.95657 6.67591 2.94637 6.35949 3.13523 6.15803Z" fill="currentColor" fill-rule="evenodd" clip-rule="evenodd"></path></svg>`,
|
||||
onclick: ()=>{
|
||||
if(info.images.length <= 2) return
|
||||
_this.imgList.classList.remove("no-transition")
|
||||
if(_this.imgCurrent == 0){
|
||||
_this.imgCurrent = (info.images.length/2)-1
|
||||
this.slides[this.slides.length-1].style.transform = 'translate3d(' + (-660 * (this.imgCurrent+1)).toString()+'px, 0px, 0px)'
|
||||
this.slides[this.slides.length-2].style.transform = 'translate3d(' + (-660 * (this.imgCurrent+1)).toString()+'px, 0px, 0px)'
|
||||
_this.imgList.style.transform = 'translate3d(660px, 0px, 0px)'
|
||||
setTimeout(_=>{
|
||||
this.slides[this.slides.length-1].style.transform = 'translate3d(0px, 0px, 0px)'
|
||||
this.slides[this.slides.length-2].style.transform = 'translate3d(0px, 0px, 0px)'
|
||||
_this.imgDistance = (-660 * this.imgCurrent).toString()
|
||||
_this.imgList.style.transform = 'translate3d(' + _this.imgDistance +'px, 0px, 0px)'
|
||||
_this.imgList.classList.add("no-transition")
|
||||
},500)
|
||||
}
|
||||
else {
|
||||
_this.imgCurrent = _this.imgCurrent-1
|
||||
_this.imgDistance = (-660 * this.imgCurrent).toString()
|
||||
_this.imgList.style.transform = 'translate3d(' + _this.imgDistance +'px, 0px, 0px)'
|
||||
}
|
||||
}
|
||||
})
|
||||
this.slideRightButton = $el("button.right",{
|
||||
parent: this.imgWrapper,
|
||||
style:{
|
||||
display:info.images.length <= 2 ? 'none' : 'block'
|
||||
},
|
||||
innerHTML:`<svg viewBox="0 0 15 15" fill="none" xmlns="http://www.w3.org/2000/svg" width="16" height="16" style="transform: rotate(-90deg);"><path d="M3.13523 6.15803C3.3241 5.95657 3.64052 5.94637 3.84197 6.13523L7.5 9.56464L11.158 6.13523C11.3595 5.94637 11.6759 5.95657 11.8648 6.15803C12.0536 6.35949 12.0434 6.67591 11.842 6.86477L7.84197 10.6148C7.64964 10.7951 7.35036 10.7951 7.15803 10.6148L3.15803 6.86477C2.95657 6.67591 2.94637 6.35949 3.13523 6.15803Z" fill="currentColor" fill-rule="evenodd" clip-rule="evenodd"></path></svg>`,
|
||||
onclick: ()=>{
|
||||
if(info.images.length <= 2) return
|
||||
_this.imgList.classList.remove("no-transition")
|
||||
|
||||
if( _this.imgCurrent >= (info.images.length/2)-1){
|
||||
_this.imgCurrent = 0
|
||||
const max = info.images.length/2
|
||||
this.slides[0].style.transform = 'translate3d(' + (660 * max).toString()+'px, 0px, 0px)'
|
||||
this.slides[1].style.transform = 'translate3d(' + (660 * max).toString()+'px, 0px, 0px)'
|
||||
_this.imgList.style.transform = 'translate3d(' + (-660 * max).toString()+'px, 0px, 0px)'
|
||||
setTimeout(_=>{
|
||||
this.slides[0].style.transform = 'translate3d(0px, 0px, 0px)'
|
||||
this.slides[1].style.transform = 'translate3d(0px, 0px, 0px)'
|
||||
_this.imgDistance = (-660 * this.imgCurrent).toString()
|
||||
_this.imgList.style.transform = 'translate3d(' + _this.imgDistance +'px, 0px, 0px)'
|
||||
_this.imgList.classList.add("no-transition")
|
||||
},500)
|
||||
}
|
||||
else {
|
||||
_this.imgCurrent = _this.imgCurrent+1
|
||||
_this.imgDistance = (-660 * this.imgCurrent).toString()
|
||||
_this.imgList.style.transform = 'translate3d(' + _this.imgDistance +'px, 0px, 0px)'
|
||||
}
|
||||
|
||||
}
|
||||
})
|
||||
|
||||
}
|
||||
|
||||
if(info.description){
|
||||
$el("div", {
|
||||
parent: this.content,
|
||||
innerHTML: info.description,
|
||||
style: {
|
||||
marginTop: "10px",
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
return info;
|
||||
})
|
||||
.catch((err) => {
|
||||
this.imgWrapper.style.display = 'none'
|
||||
content.textContent = "⚠️ " + err.message;
|
||||
})
|
||||
.finally(_=>{
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
export class CheckpointInfoDialog extends ModelInfoDialog {
|
||||
async addInfo() {
|
||||
// super.addInfo();
|
||||
await this.addCivitaiInfo();
|
||||
}
|
||||
}
|
||||
|
||||
const MAX_TAGS = 500
|
||||
export class LoraInfoDialog extends ModelInfoDialog {
|
||||
getTagFrequency() {
|
||||
if (!this.metadata.ss_tag_frequency) return [];
|
||||
|
||||
const datasets = JSON.parse(this.metadata.ss_tag_frequency);
|
||||
const tags = {};
|
||||
for (const setName in datasets) {
|
||||
const set = datasets[setName];
|
||||
for (const t in set) {
|
||||
if (t in tags) {
|
||||
tags[t] += set[t];
|
||||
} else {
|
||||
tags[t] = set[t];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return Object.entries(tags).sort((a, b) => b[1] - a[1]);
|
||||
}
|
||||
|
||||
getResolutions() {
|
||||
let res = [];
|
||||
if (this.metadata.ss_bucket_info) {
|
||||
const parsed = JSON.parse(this.metadata.ss_bucket_info);
|
||||
if (parsed?.buckets) {
|
||||
for (const { resolution, count } of Object.values(parsed.buckets)) {
|
||||
res.push([count, `${resolution.join("x")} * ${count}`]);
|
||||
}
|
||||
}
|
||||
}
|
||||
res = res.sort((a, b) => b[0] - a[0]).map((a) => a[1]);
|
||||
let r = this.metadata.ss_resolution;
|
||||
if (r) {
|
||||
const s = r.split(",");
|
||||
const w = s[0].replace("(", "");
|
||||
const h = s[1].replace(")", "");
|
||||
res.push(`${w.trim()}x${h.trim()} (Base res)`);
|
||||
} else if ((r = this.metadata["modelspec.resolution"])) {
|
||||
res.push(r + " (Base res");
|
||||
}
|
||||
if (!res.length) {
|
||||
res.push("⚠️ Unknown");
|
||||
}
|
||||
return res;
|
||||
}
|
||||
|
||||
getTagList(tags) {
|
||||
return tags.map((t) =>
|
||||
$el(
|
||||
"li.easyuse-model-tag",
|
||||
{
|
||||
dataset: {
|
||||
tag: t[0],
|
||||
},
|
||||
$: (el) => {
|
||||
el.onclick = () => {
|
||||
el.classList.toggle("easyuse-model-tag--selected");
|
||||
};
|
||||
},
|
||||
},
|
||||
[
|
||||
$el("p", {
|
||||
textContent: t[0],
|
||||
}),
|
||||
$el("span", {
|
||||
textContent: t[1],
|
||||
}),
|
||||
]
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
addTags() {
|
||||
let tags = this.getTagFrequency();
|
||||
let hasMore;
|
||||
if (tags?.length) {
|
||||
const c = tags.length;
|
||||
let list;
|
||||
if (c > MAX_TAGS) {
|
||||
tags = tags.slice(0, MAX_TAGS);
|
||||
hasMore = $el("p", [
|
||||
$el("span", { textContent: `⚠️ Only showing first ${MAX_TAGS} tags ` }),
|
||||
$el("a", {
|
||||
href: "#",
|
||||
textContent: `Show all ${c}`,
|
||||
onclick: () => {
|
||||
list.replaceChildren(...this.getTagList(this.getTagFrequency()));
|
||||
hasMore.remove();
|
||||
},
|
||||
}),
|
||||
]);
|
||||
}
|
||||
list = $el("ol.easyuse-model-tags-list", this.getTagList(tags));
|
||||
this.tags = $el("div", [list]);
|
||||
} else {
|
||||
this.tags = $el("p", { textContent: "⚠️ No tag frequency metadata found" });
|
||||
}
|
||||
|
||||
this.content.append(this.tags);
|
||||
|
||||
if (hasMore) {
|
||||
this.content.append(hasMore);
|
||||
}
|
||||
}
|
||||
|
||||
async addInfo() {
|
||||
// this.addInfoEntry("Name", this.metadata.ss_output_name || "⚠️ Unknown");
|
||||
// this.addInfoEntry("Base Model", this.metadata.ss_sd_model_name || "⚠️ Unknown");
|
||||
// this.addInfoEntry("Clip Skip", this.metadata.ss_clip_skip || "⚠️ Unknown");
|
||||
//
|
||||
// this.addInfoEntry(
|
||||
// "Resolution",
|
||||
// $el(
|
||||
// "select",
|
||||
// this.getResolutions().map((r) => $el("option", { textContent: r }))
|
||||
// )
|
||||
// );
|
||||
|
||||
// super.addInfo();
|
||||
const p = this.addCivitaiInfo();
|
||||
this.addTags();
|
||||
|
||||
const info = await p;
|
||||
if (info) {
|
||||
// $el(
|
||||
// "p",
|
||||
// {
|
||||
// parent: this.content,
|
||||
// textContent: "Trained Words: ",
|
||||
// },
|
||||
// [
|
||||
// $el("pre", {
|
||||
// textContent: info.trainedWords.join(", "),
|
||||
// style: {
|
||||
// whiteSpace: "pre-wrap",
|
||||
// margin: "10px 0",
|
||||
// background: "#222",
|
||||
// padding: "5px",
|
||||
// borderRadius: "5px",
|
||||
// maxHeight: "250px",
|
||||
// overflow: "auto",
|
||||
// },
|
||||
// }),
|
||||
// ]
|
||||
// );
|
||||
$el("div", {
|
||||
parent: this.content,
|
||||
innerHTML: info.description,
|
||||
style: {
|
||||
maxHeight: "250px",
|
||||
overflow: "auto",
|
||||
},
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
createButtons() {
|
||||
const btns = super.createButtons();
|
||||
|
||||
function copyTags(e, tags) {
|
||||
const textarea = $el("textarea", {
|
||||
parent: document.body,
|
||||
style: {
|
||||
position: "fixed",
|
||||
},
|
||||
textContent: tags.map((el) => el.dataset.tag).join(", "),
|
||||
});
|
||||
textarea.select();
|
||||
try {
|
||||
document.execCommand("copy");
|
||||
if (!e.target.dataset.text) {
|
||||
e.target.dataset.text = e.target.textContent;
|
||||
}
|
||||
e.target.textContent = "Copied " + tags.length + " tags";
|
||||
setTimeout(() => {
|
||||
e.target.textContent = e.target.dataset.text;
|
||||
}, 1000);
|
||||
} catch (ex) {
|
||||
prompt("Copy to clipboard: Ctrl+C, Enter", text);
|
||||
} finally {
|
||||
document.body.removeChild(textarea);
|
||||
}
|
||||
}
|
||||
|
||||
btns.unshift(
|
||||
$el("button", {
|
||||
type: "button",
|
||||
textContent: "Copy Selected",
|
||||
onclick: (e) => {
|
||||
copyTags(e, [...this.tags.querySelectorAll(".easyuse-model-tag--selected")]);
|
||||
},
|
||||
}),
|
||||
$el("button", {
|
||||
type: "button",
|
||||
textContent: "Copy All",
|
||||
onclick: (e) => {
|
||||
copyTags(e, [...this.tags.querySelectorAll(".easyuse-model-tag")]);
|
||||
},
|
||||
})
|
||||
);
|
||||
|
||||
return btns;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
import {sleep} from "./utils.js";
|
||||
import {$t} from "./i18n.js";
|
||||
|
||||
class Toast{
|
||||
|
||||
constructor() {
|
||||
this.info_icon = `<svg focusable="false" data-icon="info-circle" width="1em" height="1em" fill="currentColor" aria-hidden="true" viewBox="64 64 896 896"><path d="M512 64C264.6 64 64 264.6 64 512s200.6 448 448 448 448-200.6 448-448S759.4 64 512 64zm32 664c0 4.4-3.6 8-8 8h-48c-4.4 0-8-3.6-8-8V456c0-4.4 3.6-8 8-8h48c4.4 0 8 3.6 8 8v272zm-32-344a48.01 48.01 0 010-96 48.01 48.01 0 010 96z"></path></svg>`
|
||||
this.success_icon = `<svg focusable="false" data-icon="check-circle" width="1em" height="1em" fill="currentColor" aria-hidden="true" viewBox="64 64 896 896"><path d="M512 64C264.6 64 64 264.6 64 512s200.6 448 448 448 448-200.6 448-448S759.4 64 512 64zm193.5 301.7l-210.6 292a31.8 31.8 0 01-51.7 0L318.5 484.9c-3.8-5.3 0-12.7 6.5-12.7h46.9c10.2 0 19.9 4.9 25.9 13.3l71.2 98.8 157.2-218c6-8.3 15.6-13.3 25.9-13.3H699c6.5 0 10.3 7.4 6.5 12.7z"></path></svg>`
|
||||
this.error_icon = `<svg focusable="false" data-icon="close-circle" width="1em" height="1em" fill="currentColor" aria-hidden="true" fill-rule="evenodd" viewBox="64 64 896 896"><path d="M512 64c247.4 0 448 200.6 448 448S759.4 960 512 960 64 759.4 64 512 264.6 64 512 64zm127.98 274.82h-.04l-.08.06L512 466.75 384.14 338.88c-.04-.05-.06-.06-.08-.06a.12.12 0 00-.07 0c-.03 0-.05.01-.09.05l-45.02 45.02a.2.2 0 00-.05.09.12.12 0 000 .07v.02a.27.27 0 00.06.06L466.75 512 338.88 639.86c-.05.04-.06.06-.06.08a.12.12 0 000 .07c0 .03.01.05.05.09l45.02 45.02a.2.2 0 00.09.05.12.12 0 00.07 0c.02 0 .04-.01.08-.05L512 557.25l127.86 127.87c.04.04.06.05.08.05a.12.12 0 00.07 0c.03 0 .05-.01.09-.05l45.02-45.02a.2.2 0 00.05-.09.12.12 0 000-.07v-.02a.27.27 0 00-.05-.06L557.25 512l127.87-127.86c.04-.04.05-.06.05-.08a.12.12 0 000-.07c0-.03-.01-.05-.05-.09l-45.02-45.02a.2.2 0 00-.09-.05.12.12 0 00-.07 0z"></path></svg>`
|
||||
this.warn_icon = `<svg focusable="false" data-icon="exclamation-circle" width="1em" height="1em" fill="currentColor" aria-hidden="true" viewBox="64 64 896 896"><path d="M512 64C264.6 64 64 264.6 64 512s200.6 448 448 448 448-200.6 448-448S759.4 64 512 64zm-32 232c0-4.4 3.6-8 8-8h48c4.4 0 8 3.6 8 8v272c0 4.4-3.6 8-8 8h-48c-4.4 0-8-3.6-8-8V296zm32 440a48.01 48.01 0 010-96 48.01 48.01 0 010 96z"></path></svg>`
|
||||
this.loading_icon = `<svg focusable="false" data-icon="loading" width="1em" height="1em" fill="currentColor" aria-hidden="true" viewBox="0 0 1024 1024"><path d="M988 548c-19.9 0-36-16.1-36-36 0-59.4-11.6-117-34.6-171.3a440.45 440.45 0 00-94.3-139.9 437.71 437.71 0 00-139.9-94.3C629 83.6 571.4 72 512 72c-19.9 0-36-16.1-36-36s16.1-36 36-36c69.1 0 136.2 13.5 199.3 40.3C772.3 66 827 103 874 150c47 47 83.9 101.8 109.7 162.7 26.7 63.1 40.2 130.2 40.2 199.3.1 19.9-16 36-35.9 36z"></path></svg>`
|
||||
}
|
||||
|
||||
async showToast(data){
|
||||
let container = document.querySelector(".easyuse-toast-container");
|
||||
if (!container) {
|
||||
container = document.createElement("div");
|
||||
container.classList.add("easyuse-toast-container");
|
||||
document.body.appendChild(container);
|
||||
}
|
||||
await this.hideToast(data.id);
|
||||
const toastContainer = document.createElement("div");
|
||||
const content = document.createElement("span");
|
||||
content.innerHTML = data.content;
|
||||
toastContainer.appendChild(content);
|
||||
for (let a = 0; a < (data.actions || []).length; a++) {
|
||||
const action = data.actions[a];
|
||||
if (a > 0) {
|
||||
const sep = document.createElement("span");
|
||||
sep.innerHTML = " | ";
|
||||
toastContainer.appendChild(sep);
|
||||
}
|
||||
const actionEl = document.createElement("a");
|
||||
actionEl.innerText = action.label;
|
||||
if (action.href) {
|
||||
actionEl.target = "_blank";
|
||||
actionEl.href = action.href;
|
||||
}
|
||||
if (action.callback) {
|
||||
actionEl.onclick = (e) => {
|
||||
return action.callback(e);
|
||||
};
|
||||
}
|
||||
toastContainer.appendChild(actionEl);
|
||||
}
|
||||
const animContainer = document.createElement("div");
|
||||
animContainer.setAttribute("toast-id", data.id);
|
||||
animContainer.appendChild(toastContainer);
|
||||
container.appendChild(animContainer);
|
||||
await sleep(64);
|
||||
animContainer.style.marginTop = `-${animContainer.offsetHeight}px`;
|
||||
await sleep(64);
|
||||
animContainer.classList.add("-show");
|
||||
if (data.duration) {
|
||||
await sleep(data.duration);
|
||||
this.hideToast(data.id);
|
||||
}
|
||||
}
|
||||
async hideToast(id) {
|
||||
const msg = document.querySelector(`.easyuse-toast-container > [toast-id="${id}"]`);
|
||||
if (msg === null || msg === void 0 ? void 0 : msg.classList.contains("-show")) {
|
||||
msg.classList.remove("-show");
|
||||
await sleep(750);
|
||||
}
|
||||
msg && msg.remove();
|
||||
}
|
||||
async clearAllMessages() {
|
||||
let container = document.querySelector(".easyuse-toast-container");
|
||||
container && (container.innerHTML = "");
|
||||
}
|
||||
|
||||
async copyright(duration = 5000, actions = []) {
|
||||
this.showToast({
|
||||
id: `toast-info`,
|
||||
content: `${this.info_icon} ${$t('Workflow created by')} <a href="https://github.com/yolain/">Yolain</a> , ${$t('Watch more video content')} <a href="https://space.bilibili.com/1840885116">B站乱乱呀</a>`,
|
||||
duration,
|
||||
actions
|
||||
});
|
||||
}
|
||||
async info(content, duration = 3000, actions = []) {
|
||||
this.showToast({
|
||||
id: `toast-info`,
|
||||
content: `${this.info_icon} ${content}`,
|
||||
duration,
|
||||
actions
|
||||
});
|
||||
}
|
||||
async success(content, duration = 3000, actions = []) {
|
||||
this.showToast({
|
||||
id: `toast-success`,
|
||||
content: `${this.success_icon} ${content}`,
|
||||
duration,
|
||||
actions
|
||||
});
|
||||
}
|
||||
async error(content, duration = 3000, actions = []) {
|
||||
this.showToast({
|
||||
id: `toast-error`,
|
||||
content: `${this.error_icon} ${content}`,
|
||||
duration,
|
||||
actions
|
||||
});
|
||||
}
|
||||
async warn(content, duration = 3000, actions = []) {
|
||||
this.showToast({
|
||||
id: `toast-warn`,
|
||||
content: `${this.warn_icon} ${content}`,
|
||||
duration,
|
||||
actions
|
||||
});
|
||||
}
|
||||
async showLoading(content, duration = 0, actions = []) {
|
||||
this.showToast({
|
||||
id: `toast-loading`,
|
||||
content: `${this.loading_icon} ${content}`,
|
||||
duration,
|
||||
actions
|
||||
});
|
||||
}
|
||||
|
||||
async hideLoading() {
|
||||
this.hideToast("toast-loading");
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
export const toast = new Toast();
|
||||
+140
-1
@@ -1,3 +1,10 @@
|
||||
export function sleep(ms = 100, value) {
|
||||
return new Promise((resolve) => {
|
||||
setTimeout(() => {
|
||||
resolve(value);
|
||||
}, ms);
|
||||
});
|
||||
}
|
||||
export function addPreconnect(href, crossorigin=false){
|
||||
const preconnect = document.createElement("link");
|
||||
preconnect.rel = 'preconnect'
|
||||
@@ -5,7 +12,6 @@ export function addPreconnect(href, crossorigin=false){
|
||||
if(crossorigin) preconnect.crossorigin = ''
|
||||
document.head.appendChild(preconnect);
|
||||
}
|
||||
|
||||
export function addCss(href, base=true) {
|
||||
const link = document.createElement("link");
|
||||
link.rel = "stylesheet";
|
||||
@@ -38,4 +44,137 @@ export function deepEqual(obj1, obj2) {
|
||||
export function getLocale(){
|
||||
const locale = localStorage['AGL.Locale'] || localStorage['Comfy.Settings.AGL.Locale'] || 'en-US'
|
||||
return locale
|
||||
}
|
||||
|
||||
export function spliceExtension(fileName){
|
||||
return fileName.substring(0,fileName.lastIndexOf('.'))
|
||||
}
|
||||
export function getExtension(fileName){
|
||||
return fileName.substring(fileName.lastIndexOf('.') + 1)
|
||||
}
|
||||
|
||||
export function formatTime(time, format) {
|
||||
time = typeof (time) === "number" ? time : (time instanceof Date ? time.getTime() : parseInt(time));
|
||||
if (isNaN(time)) return null;
|
||||
if (typeof (format) !== 'string' || !format) format = 'yyyy-MM-dd hh:mm:ss';
|
||||
let _time = new Date(time);
|
||||
time = _time.toString().split(/[\s\:]/g).slice(0, -2);
|
||||
time[1] = ['01', '02', '03', '04', '05', '06', '07', '08', '09', '10', '11', '12'][_time.getMonth()];
|
||||
let _mapping = {
|
||||
MM: 1,
|
||||
dd: 2,
|
||||
yyyy: 3,
|
||||
hh: 4,
|
||||
mm: 5,
|
||||
ss: 6
|
||||
};
|
||||
return format.replace(/([Mmdhs]|y{2})\1/g, (key) => time[_mapping[key]]);
|
||||
}
|
||||
|
||||
|
||||
let origProps = {};
|
||||
export const findWidgetByName = (node, name) => node.widgets.find((w) => w.name === name);
|
||||
|
||||
export const doesInputWithNameExist = (node, name) => node.inputs ? node.inputs.some((input) => input.name === name) : false;
|
||||
|
||||
export function updateNodeHeight(node) {node.setSize([node.size[0], node.computeSize()[1]]);}
|
||||
|
||||
export function toggleWidget(node, widget, show = false, suffix = "") {
|
||||
if (!widget || doesInputWithNameExist(node, widget.name)) return;
|
||||
if (!origProps[widget.name]) {
|
||||
origProps[widget.name] = { origType: widget.type, origComputeSize: widget.computeSize };
|
||||
}
|
||||
const origSize = node.size;
|
||||
|
||||
widget.type = show ? origProps[widget.name].origType : "easyHidden" + suffix;
|
||||
widget.computeSize = show ? origProps[widget.name].origComputeSize : () => [0, -4];
|
||||
|
||||
widget.linkedWidgets?.forEach(w => toggleWidget(node, w, ":" + widget.name, show));
|
||||
|
||||
const height = show ? Math.max(node.computeSize()[1], origSize[1]) : node.size[1];
|
||||
node.setSize([node.size[0], height]);
|
||||
}
|
||||
|
||||
export function isLocalNetwork(ip) {
|
||||
const localNetworkRanges = [
|
||||
'192.168.',
|
||||
'10.',
|
||||
'127.',
|
||||
/^172\.((1[6-9]|2[0-9]|3[0-1])\.)/
|
||||
];
|
||||
|
||||
return localNetworkRanges.some(range => {
|
||||
if (typeof range === 'string') {
|
||||
return ip.startsWith(range);
|
||||
} else {
|
||||
return range.test(ip);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
|
||||
/**
|
||||
* accAdd 高精度加法
|
||||
* @since 1.0.10
|
||||
* @param {Number} arg1
|
||||
* @param {Number} arg2
|
||||
* @return {Number}
|
||||
*/
|
||||
export function accAdd(arg1, arg2) {
|
||||
let r1, r2, s1, s2,max;
|
||||
s1 = typeof arg1 == 'string' ? arg1 : arg1.toString()
|
||||
s2 = typeof arg2 == 'string' ? arg2 : arg2.toString()
|
||||
try { r1 = s1.split(".")[1].length } catch (e) { r1 = 0 }
|
||||
try { r2 = s2.split(".")[1].length } catch (e) { r2 = 0 }
|
||||
max = Math.pow(10, Math.max(r1, r2))
|
||||
return (arg1 * max + arg2 * max) / max
|
||||
}
|
||||
/**
|
||||
* accSub 高精度减法
|
||||
* @since 1.0.10
|
||||
* @param {Number} arg1
|
||||
* @param {Number} arg2
|
||||
* @return {Number}
|
||||
*/
|
||||
export function accSub(arg1, arg2) {
|
||||
let r1, r2, max, min,s1,s2;
|
||||
s1 = typeof arg1 == 'string' ? arg1 : arg1.toString()
|
||||
s2 = typeof arg2 == 'string' ? arg2 : arg2.toString()
|
||||
try { r1 = s1.split(".")[1].length } catch (e) { r1 = 0 }
|
||||
try { r2 = s2.split(".")[1].length } catch (e) { r2 = 0 }
|
||||
max = Math.pow(10, Math.max(r1, r2));
|
||||
//动态控制精度长度
|
||||
min = (r1 >= r2) ? r1 : r2;
|
||||
return ((arg1 * max - arg2 * max) / max).toFixed(min)
|
||||
}
|
||||
/**
|
||||
* accMul 高精度乘法
|
||||
* @since 1.0.10
|
||||
* @param {Number} arg1
|
||||
* @param {Number} arg2
|
||||
* @return {Number}
|
||||
*/
|
||||
export function accMul(arg1, arg2) {
|
||||
let max = 0, s1 = typeof arg1 == 'string' ? arg1 : arg1.toString(), s2 = typeof arg2 == 'string' ? arg2 : arg2.toString();
|
||||
try { max += s1.split(".")[1].length } catch (e) { }
|
||||
try { max += s2.split(".")[1].length } catch (e) { }
|
||||
return Number(s1.replace(".", "")) * Number(s2.replace(".", "")) / Math.pow(10, max)
|
||||
}
|
||||
/**
|
||||
* accDiv 高精度除法
|
||||
* @since 1.0.10
|
||||
* @param {Number} arg1
|
||||
* @param {Number} arg2
|
||||
* @return {Number}
|
||||
*/
|
||||
export function accDiv(arg1, arg2) {
|
||||
let t1 = 0, t2 = 0, r1, r2,s1 = typeof arg1 == 'string' ? arg1 : arg1.toString(), s2 = typeof arg2 == 'string' ? arg2 : arg2.toString();
|
||||
try { t1 = s1.toString().split(".")[1].length } catch (e) { }
|
||||
try { t2 = s2.toString().split(".")[1].length } catch (e) { }
|
||||
r1 = Number(s1.toString().replace(".", ""))
|
||||
r2 = Number(s2.toString().replace(".", ""))
|
||||
return (r1 / r2) * Math.pow(10, t2 - t1)
|
||||
}
|
||||
Number.prototype.div = function (arg) {
|
||||
return accDiv(this, arg);
|
||||
}
|
||||
+500
-234
@@ -1,9 +1,419 @@
|
||||
import { app } from "/scripts/app.js";
|
||||
import {deepEqual,addCss} from "../common/utils.js";
|
||||
import { api } from "../../../../scripts/api.js";
|
||||
import { app } from "../../../../scripts/app.js";
|
||||
import {deepEqual, addCss, isLocalNetwork} from "../common/utils.js";
|
||||
import {quesitonIcon, rocketIcon, groupIcon, rebootIcon, closeIcon} from "../common/icon.js";
|
||||
import {$t} from '../common/i18n.js';
|
||||
import {toast} from "../common/toast.js";
|
||||
import {$el, ComfyDialog} from "../../../../scripts/ui.js";
|
||||
|
||||
|
||||
addCss('css/index.css')
|
||||
|
||||
api.addEventListener("easyuse-toast",event=>{
|
||||
const content = event.detail.content
|
||||
const type = event.detail.type
|
||||
const duration = event.detail.duration
|
||||
if(!type){
|
||||
toast.info(content, duration)
|
||||
}
|
||||
else{
|
||||
toast.showToast({
|
||||
id: `toast-${type}`,
|
||||
content: `${toast[type+"_icon"]} ${content}`,
|
||||
duration: duration || 3000,
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
let draggerEl = null
|
||||
let isGroupMapcanMove = true
|
||||
function createGroupMap(){
|
||||
let div = document.querySelector('#easyuse_groups_map')
|
||||
if(div){
|
||||
div.style.display = div.style.display == 'none' ? 'flex' : 'none'
|
||||
return
|
||||
}
|
||||
let groups = app.canvas.graph._groups
|
||||
let nodes = app.canvas.graph._nodes
|
||||
let old_nodes = groups.length
|
||||
div = document.createElement('div')
|
||||
div.id = 'easyuse_groups_map'
|
||||
div.innerHTML = ''
|
||||
let btn = document.createElement('div')
|
||||
btn.style = `display: flex;
|
||||
width: calc(100% - 8px);
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
padding: 0 6px;
|
||||
height: 44px;`
|
||||
let hideBtn = $el('button.closeBtn',{
|
||||
innerHTML:closeIcon,
|
||||
onclick:_=>div.style.display = 'none'
|
||||
})
|
||||
let textB = document.createElement('p')
|
||||
btn.appendChild(textB)
|
||||
btn.appendChild(hideBtn)
|
||||
textB.style.fontSize = '11px'
|
||||
textB.innerHTML = `<b>${$t('Groups Map')} (EasyUse)</b>`
|
||||
div.appendChild(btn)
|
||||
|
||||
div.addEventListener('mousedown', function (e) {
|
||||
var startX = e.clientX
|
||||
var startY = e.clientY
|
||||
var offsetX = div.offsetLeft
|
||||
var offsetY = div.offsetTop
|
||||
|
||||
function moveBox (e) {
|
||||
var newX = e.clientX
|
||||
var newY = e.clientY
|
||||
var deltaX = newX - startX
|
||||
var deltaY = newY - startY
|
||||
div.style.left = offsetX + deltaX + 'px'
|
||||
div.style.top = offsetY + deltaY + 'px'
|
||||
}
|
||||
|
||||
function stopMoving () {
|
||||
document.removeEventListener('mousemove', moveBox)
|
||||
document.removeEventListener('mouseup', stopMoving)
|
||||
}
|
||||
|
||||
if(isGroupMapcanMove){
|
||||
document.addEventListener('mousemove', moveBox)
|
||||
document.addEventListener('mouseup', stopMoving)
|
||||
}
|
||||
})
|
||||
|
||||
function updateGroups(groups, groupsDiv, autoSortDiv){
|
||||
if(groups.length>0){
|
||||
autoSortDiv.style.display = 'block'
|
||||
}else autoSortDiv.style.display = 'none'
|
||||
for (let index in groups) {
|
||||
const group = groups[index]
|
||||
const title = group.title
|
||||
const show_text = $t('Always')
|
||||
const hide_text = $t('Bypass')
|
||||
const mute_text = $t('Never')
|
||||
let group_item = document.createElement('div')
|
||||
let group_item_style = `justify-content: space-between;display:flex;background-color: var(--comfy-input-bg);border-radius: 5px;border:1px solid var(--border-color);margin-top:5px;`
|
||||
group_item.addEventListener("mouseover",event=>{
|
||||
event.preventDefault()
|
||||
group_item.style = group_item_style + "filter:brightness(1.2);"
|
||||
})
|
||||
group_item.addEventListener("mouseleave",event=>{
|
||||
event.preventDefault()
|
||||
group_item.style = group_item_style + "filter:brightness(1);"
|
||||
})
|
||||
group_item.addEventListener("dragstart",e=>{
|
||||
draggerEl = e.currentTarget;
|
||||
e.currentTarget.style.opacity = "0.6";
|
||||
e.currentTarget.style.border = "1px dashed yellow";
|
||||
e.dataTransfer.effectAllowed = 'move';
|
||||
e.dataTransfer.setDragImage(emptyImg, 0, 0);
|
||||
})
|
||||
group_item.addEventListener("dragend",e=>{
|
||||
e.target.style.opacity = "1";
|
||||
e.currentTarget.style.border = "1px dashed transparent";
|
||||
e.currentTarget.removeAttribute("draggable");
|
||||
document.querySelectorAll('.easyuse-group-item').forEach((el,i) => {
|
||||
var prev_i = el.dataset.id;
|
||||
if (el == draggerEl && prev_i != i ) {
|
||||
groups.splice(i, 0, groups.splice(prev_i, 1)[0]);
|
||||
}
|
||||
el.dataset.id = i;
|
||||
});
|
||||
isGroupMapcanMove = true
|
||||
})
|
||||
group_item.addEventListener("dragover",e=>{
|
||||
e.preventDefault();
|
||||
if (e.currentTarget == draggerEl) return;
|
||||
let rect = e.currentTarget.getBoundingClientRect();
|
||||
if (e.clientY > rect.top + rect.height / 2) {
|
||||
e.currentTarget.parentNode.insertBefore(draggerEl, e.currentTarget.nextSibling);
|
||||
} else {
|
||||
e.currentTarget.parentNode.insertBefore(draggerEl, e.currentTarget);
|
||||
}
|
||||
isGroupMapcanMove = true
|
||||
})
|
||||
|
||||
|
||||
group_item.setAttribute('data-id',index)
|
||||
group_item.className = 'easyuse-group-item'
|
||||
group_item.style = group_item_style
|
||||
// 标题
|
||||
let text_group_title = document.createElement('div')
|
||||
text_group_title.style = `flex:1;font-size:12px;color:var(--input-text);padding:4px;white-space: nowrap;overflow: hidden;text-overflow: ellipsis;cursor:pointer`
|
||||
text_group_title.innerHTML = `${title}`
|
||||
text_group_title.addEventListener('mousedown',e=>{
|
||||
isGroupMapcanMove = false
|
||||
e.currentTarget.parentNode.draggable = 'true';
|
||||
})
|
||||
text_group_title.addEventListener('mouseleave',e=>{
|
||||
setTimeout(_=>{
|
||||
isGroupMapcanMove = true
|
||||
},150)
|
||||
})
|
||||
group_item.append(text_group_title)
|
||||
// 按钮组
|
||||
let buttons = document.createElement('div')
|
||||
group.recomputeInsideNodes();
|
||||
const nodesInGroup = group._nodes;
|
||||
let isGroupShow = nodesInGroup && nodesInGroup.length>0 && nodesInGroup[0].mode == 0
|
||||
let isGroupMute = nodesInGroup && nodesInGroup.length>0 && nodesInGroup[0].mode == 2
|
||||
let go_btn = document.createElement('button')
|
||||
go_btn.style = "margin-right:6px;cursor:pointer;font-size:10px;padding:2px 4px;color:var(--input-text);background-color: var(--comfy-input-bg);border: 1px solid var(--border-color);border-radius:4px;"
|
||||
go_btn.innerText = "Go"
|
||||
go_btn.addEventListener('click', () => {
|
||||
app.canvas.ds.offset[0] = -group.pos[0] - group.size[0] * 0.5 + (app.canvas.canvas.width * 0.5) / app.canvas.ds.scale;
|
||||
app.canvas.ds.offset[1] = -group.pos[1] - group.size[1] * 0.5 + (app.canvas.canvas.height * 0.5) / app.canvas.ds.scale;
|
||||
app.canvas.setDirty(true, true);
|
||||
app.canvas.setZoom(1)
|
||||
})
|
||||
buttons.append(go_btn)
|
||||
let see_btn = document.createElement('button')
|
||||
let defaultStyle = `cursor:pointer;font-size:10px;;padding:2px;border: 1px solid var(--border-color);border-radius:4px;width:36px;`
|
||||
see_btn.style = isGroupMute ? `background-color:var(--error-text);color:var(--input-text);` + defaultStyle : (isGroupShow ? `background-color:#006691;color:var(--input-text);` + defaultStyle : `background-color: var(--comfy-input-bg);color:var(--descrip-text);` + defaultStyle)
|
||||
see_btn.innerText = isGroupMute ? mute_text : (isGroupShow ? show_text : hide_text)
|
||||
let pressTimer
|
||||
let firstTime =0, lastTime =0
|
||||
let isHolding = false
|
||||
see_btn.addEventListener('click', () => {
|
||||
if(isHolding){
|
||||
isHolding = false
|
||||
return
|
||||
}
|
||||
for (const node of nodesInGroup) {
|
||||
node.mode = isGroupShow ? 4 : 0;
|
||||
node.graph.change();
|
||||
}
|
||||
isGroupShow = nodesInGroup[0].mode == 0 ? true : false
|
||||
isGroupMute = nodesInGroup[0].mode == 2 ? true : false
|
||||
see_btn.style = isGroupMute ? `background-color:var(--error-text);color:var(--input-text);` + defaultStyle : (isGroupShow ? `background-color:#006691;color:var(--input-text);` + defaultStyle : `background-color: var(--comfy-input-bg);color:var(--descrip-text);` + defaultStyle)
|
||||
see_btn.innerText = isGroupMute ? mute_text : (isGroupShow ? show_text : hide_text)
|
||||
})
|
||||
see_btn.addEventListener('mousedown', () => {
|
||||
firstTime = new Date().getTime();
|
||||
clearTimeout(pressTimer);
|
||||
pressTimer = setTimeout(_=>{
|
||||
for (const node of nodesInGroup) {
|
||||
node.mode = isGroupMute ? 0 : 2;
|
||||
node.graph.change();
|
||||
}
|
||||
isGroupShow = nodesInGroup[0].mode == 0 ? true : false
|
||||
isGroupMute = nodesInGroup[0].mode == 2 ? true : false
|
||||
see_btn.style = isGroupMute ? `background-color:var(--error-text);color:var(--input-text);` + defaultStyle : (isGroupShow ? `background-color:#006691;color:var(--input-text);` + defaultStyle : `background-color: var(--comfy-input-bg);color:var(--descrip-text);` + defaultStyle)
|
||||
see_btn.innerText = isGroupMute ? mute_text : (isGroupShow ? show_text : hide_text)
|
||||
},500)
|
||||
})
|
||||
see_btn.addEventListener('mouseup', () => {
|
||||
lastTime = new Date().getTime();
|
||||
if(lastTime - firstTime > 500) isHolding = true
|
||||
clearTimeout(pressTimer);
|
||||
})
|
||||
buttons.append(see_btn)
|
||||
group_item.append(buttons)
|
||||
|
||||
groupsDiv.append(group_item)
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
let groupsDiv = document.createElement('div')
|
||||
groupsDiv.id = 'easyuse-groups-items'
|
||||
groupsDiv.style = `overflow-y: auto;max-height: 400px;height:100%;width: 100%;`
|
||||
|
||||
let autoSortDiv = document.createElement('button')
|
||||
autoSortDiv.style = `cursor:pointer;font-size:10px;padding:2px 4px;color:var(--input-text);background-color: var(--comfy-input-bg);border: 1px solid var(--border-color);border-radius:4px;`
|
||||
autoSortDiv.innerText = $t('Auto Sorting')
|
||||
autoSortDiv.addEventListener('click',e=>{
|
||||
e.preventDefault()
|
||||
groupsDiv.innerHTML = ``
|
||||
let new_groups = groups.sort((a,b)=> a['pos'][0] - b['pos'][0]).sort((a,b)=> a['pos'][1] - b['pos'][1])
|
||||
updateGroups(new_groups, groupsDiv, autoSortDiv)
|
||||
})
|
||||
|
||||
updateGroups(groups, groupsDiv, autoSortDiv)
|
||||
|
||||
div.appendChild(groupsDiv)
|
||||
|
||||
let remarkDiv = document.createElement('p')
|
||||
remarkDiv.style = `text-align:center; font-size:10px; padding:0 10px;color:var(--descrip-text)`
|
||||
remarkDiv.innerText = $t('Toggle `Show/Hide` can set mode of group, LongPress can set group nodes to never')
|
||||
div.appendChild(groupsDiv)
|
||||
div.appendChild(remarkDiv)
|
||||
div.appendChild(autoSortDiv)
|
||||
|
||||
let graphDiv = document.getElementById("graph-canvas")
|
||||
graphDiv.addEventListener('mouseover', async () => {
|
||||
groupsDiv.innerHTML = ``
|
||||
let new_groups = app.canvas.graph._groups
|
||||
updateGroups(new_groups, groupsDiv, autoSortDiv)
|
||||
old_nodes = nodes
|
||||
})
|
||||
|
||||
if (!document.querySelector('#easyuse_groups_map')){
|
||||
document.body.appendChild(div)
|
||||
}else{
|
||||
div.style.display = 'flex'
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
async function cleanup(){
|
||||
try {
|
||||
const {Running, Pending} = await api.getQueue()
|
||||
if(Running.length>0 || Pending.length>0){
|
||||
toast.error($t("Clean Failed")+ ":"+ $t("Please stop all running tasks before cleaning GPU"))
|
||||
return
|
||||
}
|
||||
api.fetchApi("/easyuse/cleangpu",{
|
||||
method:"POST"
|
||||
}).then(res=>{
|
||||
if(res.status == 200){
|
||||
toast.success($t("Clean SuccessFully"))
|
||||
}else{
|
||||
toast.error($t("Clean Failed"))
|
||||
}
|
||||
})
|
||||
|
||||
} catch (exception) {}
|
||||
}
|
||||
|
||||
|
||||
let guideDialog = null
|
||||
let isDownloading = false
|
||||
function download_model(url,local_dir){
|
||||
if(isDownloading || !url || !local_dir) return
|
||||
isDownloading = true
|
||||
let body = new FormData();
|
||||
body.append('url', url);
|
||||
body.append('local_dir', local_dir);
|
||||
api.fetchApi("/easyuse/model/download",{
|
||||
method:"POST",
|
||||
body
|
||||
}).then(res=>{
|
||||
if(res.status == 200){
|
||||
toast.success($t("Download SuccessFully"))
|
||||
}else{
|
||||
toast.error($t("Download Failed"))
|
||||
}
|
||||
isDownloading = false
|
||||
})
|
||||
|
||||
}
|
||||
class GuideDialog {
|
||||
|
||||
constructor(note, need_models){
|
||||
this.dialogDiv = null
|
||||
this.modelsDiv = null
|
||||
|
||||
if(need_models?.length>0){
|
||||
let tbody = []
|
||||
|
||||
for(let i=0;i<need_models.length;i++){
|
||||
tbody.push($el('tr',[
|
||||
$el('td',{innerHTML:need_models[i].title || need_models[i].name || ''}),
|
||||
$el('td',[
|
||||
need_models[i]['download_url'] ? $el('a',{onclick:_=>download_model(need_models[i]['download_url'],need_models[i]['local_dir']), target:"_blank", textContent:$t('Download Model')}) : '',
|
||||
need_models[i]['source_url'] ? $el('a',{href:need_models[i]['source_url'], target:"_blank", textContent:$t('Source Url')}) : '',
|
||||
need_models[i]['desciption'] ? $el('span',{textContent:need_models[i]['desciption']}) : '',
|
||||
]),
|
||||
]))
|
||||
}
|
||||
this.modelsDiv = $el('div.easyuse-guide-dialog-models.markdown-body',[
|
||||
$el('h3',{textContent:$t('Models Required')}),
|
||||
$el('table',{cellpadding:0,cellspacing:0},[
|
||||
$el('thead',[
|
||||
$el('tr',[
|
||||
$el('th',{innerHTML:$t('ModelName')}),
|
||||
$el('th',{innerHTML:$t('Description')}),
|
||||
])
|
||||
]),
|
||||
$el('tbody',tbody)
|
||||
])
|
||||
])
|
||||
}
|
||||
|
||||
this.dialogDiv = $el('div.easyuse-guide-dialog.hidden',[
|
||||
$el('div.easyuse-guide-dialog-header',[
|
||||
$el('div.easyuse-guide-dialog-top',[
|
||||
$el('div.easyuse-guide-dialog-title',{
|
||||
innerHTML:$t('Workflow Guide')
|
||||
}),
|
||||
$el('button.closeBtn',{innerHTML:closeIcon,onclick:_=>this.close()})
|
||||
]),
|
||||
|
||||
$el('div.easyuse-guide-dialog-remark',{
|
||||
innerHTML:`${$t('Workflow created by')} <a href="https://github.com/yolain/" target="_blank">Yolain</a> , ${$t('Watch more video content')} <a href="https://space.bilibili.com/1840885116" target="_blank">B站乱乱呀</a>`
|
||||
})
|
||||
]),
|
||||
$el('div.easyuse-guide-dialog-content.markdown-body',[
|
||||
$el('div.easyuse-guide-dialog-note',{
|
||||
innerHTML:note
|
||||
}),
|
||||
...this.modelsDiv ? [this.modelsDiv] : []
|
||||
])
|
||||
])
|
||||
|
||||
if(disableRenderInfo){
|
||||
this.dialogDiv.classList.add('disable-render-info')
|
||||
}
|
||||
document.body.appendChild(this.dialogDiv)
|
||||
}
|
||||
show(){
|
||||
if(this.dialogDiv) this.dialogDiv.classList.remove('hidden')
|
||||
}
|
||||
|
||||
close(){
|
||||
if(this.dialogDiv){
|
||||
this.dialogDiv.classList.add('hidden')
|
||||
}
|
||||
}
|
||||
toggle(){
|
||||
if(this.dialogDiv){
|
||||
if(this.dialogDiv.classList.contains('hidden')){
|
||||
this.show()
|
||||
}else{
|
||||
this.close()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
remove(){
|
||||
if(this.dialogDiv) document.body.removeChild(this.dialogDiv)
|
||||
}
|
||||
}
|
||||
|
||||
const getEnableToolBar = _ => app.ui.settings.getSettingValue(toolBarId, true)
|
||||
|
||||
const toolBarId = "Comfy.EasyUse.toolBar"
|
||||
|
||||
let enableToolBar = getEnableToolBar()
|
||||
let disableRenderInfo = localStorage['Comfy.Settings.Comfy.EasyUse.disableRenderInfo'] ? true : false
|
||||
export function addToolBar(app) {
|
||||
app.ui.settings.addSetting({
|
||||
id: toolBarId,
|
||||
name: $t("Enable tool bar fixed on the left-bottom (ComfyUI-Easy-Use)"),
|
||||
type: "boolean",
|
||||
defaultValue: enableToolBar,
|
||||
onChange(value) {
|
||||
enableToolBar = !!value;
|
||||
if(enableToolBar){
|
||||
showToolBar()
|
||||
}else hideToolBar()
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
let note = null
|
||||
let toolbar = null
|
||||
function showToolBar(){
|
||||
toolbar.style.display = 'flex'
|
||||
}
|
||||
function hideToolBar(){
|
||||
toolbar.style.display = 'none'
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "comfy.easyUse",
|
||||
init() {
|
||||
@@ -11,249 +421,105 @@ app.registerExtension({
|
||||
const getCanvasMenuOptions = LGraphCanvas.prototype.getCanvasMenuOptions;
|
||||
LGraphCanvas.prototype.getCanvasMenuOptions = function () {
|
||||
const options = getCanvasMenuOptions.apply(this, arguments);
|
||||
let draggerEl = null
|
||||
let isGroupMapcanMove = true
|
||||
let old_groups = []
|
||||
let emptyImg = new Image()
|
||||
emptyImg.src = "data:image/gif;base64,R0lGODlhAQABAIAAAAUEBAAAACwAAAAAAQABAAACAkQBADs=";
|
||||
|
||||
options.push(null,
|
||||
// Groups Map
|
||||
{
|
||||
content: '📜 '+ $t('Groups Map (EasyUse)'),
|
||||
content: groupIcon.replace('currentColor','var(--warning-color)') + ' '+ $t('Groups Map') + ' (EasyUse)',
|
||||
callback: async() => {
|
||||
let groups = app.canvas.graph._groups
|
||||
let nodes = app.canvas.graph._nodes
|
||||
let old_nodes = groups.length
|
||||
let div =
|
||||
document.querySelector('#easyuse_groups_map') ||
|
||||
document.createElement('div')
|
||||
div.id = 'easyuse_groups_map'
|
||||
div.innerHTML = ''
|
||||
let btn = document.createElement('div')
|
||||
btn.style = `display: flex;
|
||||
width: calc(100% - 8px);
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
padding: 0 6px;
|
||||
height: 44px;`
|
||||
let hideBtn = document.createElement('button')
|
||||
let textB = document.createElement('p')
|
||||
btn.appendChild(textB)
|
||||
btn.appendChild(hideBtn)
|
||||
textB.style.fontSize = '11px'
|
||||
textB.innerHTML = `<b>${$t('Groups Map (EasyUse)')}</b>`
|
||||
hideBtn.style = `float: right;color: var(--input-text);border-radius:6px;font-size:9px;
|
||||
background-color: var(--comfy-input-bg); border: 1px solid var(--border-color);cursor: pointer;padding: 5px;aspect-ratio: 1 / 1;`
|
||||
hideBtn.addEventListener('click', () => {
|
||||
div.style.display = 'none'
|
||||
})
|
||||
hideBtn.innerText = '❌'
|
||||
div.appendChild(btn)
|
||||
|
||||
div.addEventListener('mousedown', function (e) {
|
||||
var startX = e.clientX
|
||||
var startY = e.clientY
|
||||
var offsetX = div.offsetLeft
|
||||
var offsetY = div.offsetTop
|
||||
|
||||
function moveBox (e) {
|
||||
var newX = e.clientX
|
||||
var newY = e.clientY
|
||||
var deltaX = newX - startX
|
||||
var deltaY = newY - startY
|
||||
div.style.left = offsetX + deltaX + 'px'
|
||||
div.style.top = offsetY + deltaY + 'px'
|
||||
}
|
||||
|
||||
function stopMoving () {
|
||||
document.removeEventListener('mousemove', moveBox)
|
||||
document.removeEventListener('mouseup', stopMoving)
|
||||
}
|
||||
|
||||
if(isGroupMapcanMove){
|
||||
document.addEventListener('mousemove', moveBox)
|
||||
document.addEventListener('mouseup', stopMoving)
|
||||
}
|
||||
})
|
||||
|
||||
function updateGroups(groups, groupsDiv, autoSortDiv){
|
||||
if(groups.length>0){
|
||||
autoSortDiv.style.display = 'block'
|
||||
}else autoSortDiv.style.display = 'none'
|
||||
for (let index in groups) {
|
||||
const group = groups[index]
|
||||
const title = group.title
|
||||
const show_text = $t('Always')
|
||||
const hide_text = $t('Bypass')
|
||||
const mute_text = $t('Never')
|
||||
let group_item = document.createElement('div')
|
||||
let group_item_style = `justify-content: space-between;display:flex;background-color: var(--comfy-input-bg);border-radius: 5px;border:1px solid var(--border-color);margin-top:5px;`
|
||||
group_item.addEventListener("mouseover",event=>{
|
||||
event.preventDefault()
|
||||
group_item.style = group_item_style + "filter:brightness(1.2);"
|
||||
})
|
||||
group_item.addEventListener("mouseleave",event=>{
|
||||
event.preventDefault()
|
||||
group_item.style = group_item_style + "filter:brightness(1);"
|
||||
})
|
||||
group_item.addEventListener("dragstart",e=>{
|
||||
draggerEl = e.currentTarget;
|
||||
e.currentTarget.style.opacity = "0.6";
|
||||
e.currentTarget.style.border = "1px dashed yellow";
|
||||
e.dataTransfer.effectAllowed = 'move';
|
||||
e.dataTransfer.setDragImage(emptyImg, 0, 0);
|
||||
})
|
||||
group_item.addEventListener("dragend",e=>{
|
||||
e.target.style.opacity = "1";
|
||||
e.currentTarget.style.border = "1px dashed transparent";
|
||||
e.currentTarget.removeAttribute("draggable");
|
||||
document.querySelectorAll('.easyuse-group-item').forEach((el,i) => {
|
||||
var prev_i = el.dataset.id;
|
||||
if (el == draggerEl && prev_i != i ) {
|
||||
groups.splice(i, 0, groups.splice(prev_i, 1)[0]);
|
||||
}
|
||||
el.dataset.id = i;
|
||||
});
|
||||
isGroupMapcanMove = true
|
||||
})
|
||||
group_item.addEventListener("dragover",e=>{
|
||||
e.preventDefault();
|
||||
if (e.currentTarget == draggerEl) return;
|
||||
let rect = e.currentTarget.getBoundingClientRect();
|
||||
if (e.clientY > rect.top + rect.height / 2) {
|
||||
e.currentTarget.parentNode.insertBefore(draggerEl, e.currentTarget.nextSibling);
|
||||
} else {
|
||||
e.currentTarget.parentNode.insertBefore(draggerEl, e.currentTarget);
|
||||
}
|
||||
isGroupMapcanMove = true
|
||||
})
|
||||
|
||||
|
||||
group_item.setAttribute('data-id',index)
|
||||
group_item.className = 'easyuse-group-item'
|
||||
group_item.style = group_item_style
|
||||
// 标题
|
||||
let text_group_title = document.createElement('div')
|
||||
text_group_title.style = `flex:1;font-size:12px;color:var(--input-text);padding:4px;white-space: nowrap;overflow: hidden;text-overflow: ellipsis;cursor:pointer`
|
||||
text_group_title.innerHTML = `${title}`
|
||||
text_group_title.addEventListener('mousedown',e=>{
|
||||
isGroupMapcanMove = false
|
||||
e.currentTarget.parentNode.draggable = 'true';
|
||||
})
|
||||
text_group_title.addEventListener('mouseleave',e=>{
|
||||
setTimeout(_=>{
|
||||
isGroupMapcanMove = true
|
||||
},150)
|
||||
})
|
||||
group_item.append(text_group_title)
|
||||
// 按钮组
|
||||
let buttons = document.createElement('div')
|
||||
group.recomputeInsideNodes();
|
||||
const nodesInGroup = group._nodes;
|
||||
let isGroupShow = nodesInGroup && nodesInGroup.length>0 && nodesInGroup[0].mode == 0
|
||||
let isGroupMute = nodesInGroup && nodesInGroup.length>0 && nodesInGroup[0].mode == 2
|
||||
let go_btn = document.createElement('button')
|
||||
go_btn.style = "margin-right:6px;cursor:pointer;font-size:10px;padding:2px 4px;color:var(--input-text);background-color: var(--comfy-input-bg);border: 1px solid var(--border-color);border-radius:4px;"
|
||||
go_btn.innerText = "Go"
|
||||
go_btn.addEventListener('click', () => {
|
||||
app.canvas.ds.offset[0] = -group.pos[0] - group.size[0] * 0.5 + (app.canvas.canvas.width * 0.5) / app.canvas.ds.scale;
|
||||
app.canvas.ds.offset[1] = -group.pos[1] - group.size[1] * 0.5 + (app.canvas.canvas.height * 0.5) / app.canvas.ds.scale;
|
||||
app.canvas.setDirty(true, true);
|
||||
app.canvas.setZoom(1)
|
||||
})
|
||||
buttons.append(go_btn)
|
||||
let see_btn = document.createElement('button')
|
||||
let defaultStyle = `cursor:pointer;font-size:10px;;padding:2px;border: 1px solid var(--border-color);border-radius:4px;width:36px;`
|
||||
see_btn.style = isGroupMute ? `background-color:var(--error-text);color:var(--input-text);` + defaultStyle : (isGroupShow ? `background-color:#006691;color:var(--input-text);` + defaultStyle : `background-color: var(--comfy-input-bg);color:var(--descrip-text);` + defaultStyle)
|
||||
see_btn.innerText = isGroupMute ? mute_text : (isGroupShow ? show_text : hide_text)
|
||||
let pressTimer
|
||||
let firstTime =0, lastTime =0
|
||||
let isHolding = false
|
||||
see_btn.addEventListener('click', () => {
|
||||
if(isHolding){
|
||||
isHolding = false
|
||||
return
|
||||
}
|
||||
for (const node of nodesInGroup) {
|
||||
node.mode = isGroupShow ? 4 : 0;
|
||||
node.graph.change();
|
||||
}
|
||||
isGroupShow = nodesInGroup[0].mode == 0 ? true : false
|
||||
isGroupMute = nodesInGroup[0].mode == 2 ? true : false
|
||||
see_btn.style = isGroupMute ? `background-color:var(--error-text);color:var(--input-text);` + defaultStyle : (isGroupShow ? `background-color:#006691;color:var(--input-text);` + defaultStyle : `background-color: var(--comfy-input-bg);color:var(--descrip-text);` + defaultStyle)
|
||||
see_btn.innerText = isGroupMute ? mute_text : (isGroupShow ? show_text : hide_text)
|
||||
})
|
||||
see_btn.addEventListener('mousedown', () => {
|
||||
firstTime = new Date().getTime();
|
||||
clearTimeout(pressTimer);
|
||||
pressTimer = setTimeout(_=>{
|
||||
for (const node of nodesInGroup) {
|
||||
node.mode = isGroupMute ? 0 : 2;
|
||||
node.graph.change();
|
||||
}
|
||||
isGroupShow = nodesInGroup[0].mode == 0 ? true : false
|
||||
isGroupMute = nodesInGroup[0].mode == 2 ? true : false
|
||||
see_btn.style = isGroupMute ? `background-color:var(--error-text);color:var(--input-text);` + defaultStyle : (isGroupShow ? `background-color:#006691;color:var(--input-text);` + defaultStyle : `background-color: var(--comfy-input-bg);color:var(--descrip-text);` + defaultStyle)
|
||||
see_btn.innerText = isGroupMute ? mute_text : (isGroupShow ? show_text : hide_text)
|
||||
},500)
|
||||
})
|
||||
see_btn.addEventListener('mouseup', () => {
|
||||
lastTime = new Date().getTime();
|
||||
if(lastTime - firstTime > 500) isHolding = true
|
||||
clearTimeout(pressTimer);
|
||||
})
|
||||
buttons.append(see_btn)
|
||||
group_item.append(buttons)
|
||||
|
||||
groupsDiv.append(group_item)
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
let groupsDiv = document.createElement('div')
|
||||
groupsDiv.id = 'easyuse-groups-items'
|
||||
groupsDiv.style = `overflow-y: auto;max-height: 400px;height:100%;width: 100%;`
|
||||
|
||||
let autoSortDiv = document.createElement('button')
|
||||
autoSortDiv.style = `cursor:pointer;font-size:10px;padding:2px 4px;color:var(--input-text);background-color: var(--comfy-input-bg);border: 1px solid var(--border-color);border-radius:4px;`
|
||||
autoSortDiv.innerText = $t('Auto Sorting')
|
||||
autoSortDiv.addEventListener('click',e=>{
|
||||
e.preventDefault()
|
||||
groupsDiv.innerHTML = ``
|
||||
let new_groups = groups.sort((a,b)=> a['pos'][0] - b['pos'][0]).sort((a,b)=> a['pos'][1] - b['pos'][1])
|
||||
updateGroups(new_groups, groupsDiv, autoSortDiv)
|
||||
})
|
||||
|
||||
updateGroups(groups, groupsDiv, autoSortDiv)
|
||||
|
||||
div.appendChild(groupsDiv)
|
||||
|
||||
let remarkDiv = document.createElement('p')
|
||||
remarkDiv.style = `text-align:center; font-size:10px; padding:0 10px;color:var(--descrip-text)`
|
||||
remarkDiv.innerText = $t('Toggle `Show/Hide` can set mode of group, LongPress can set group nodes to never')
|
||||
div.appendChild(groupsDiv)
|
||||
div.appendChild(remarkDiv)
|
||||
div.appendChild(autoSortDiv)
|
||||
|
||||
let graphDiv = document.getElementById("graph-canvas")
|
||||
graphDiv.addEventListener('mouseover', async () => {
|
||||
groupsDiv.innerHTML = ``
|
||||
let new_groups = app.canvas.graph._groups
|
||||
updateGroups(new_groups, groupsDiv, autoSortDiv)
|
||||
old_nodes = nodes
|
||||
})
|
||||
|
||||
if (!document.querySelector('#easyuse_groups_map')){
|
||||
document.body.appendChild(div)
|
||||
}else{
|
||||
div.style.display = 'flex'
|
||||
}
|
||||
|
||||
createGroupMap()
|
||||
}
|
||||
},
|
||||
// Force clean ComfyUI GPU Used 强制卸载模型GPU占用
|
||||
{
|
||||
content: rocketIcon.replace('currentColor','var(--theme-color-light)') + ' '+ $t('Cleanup Of GPU Usage') + ' (EasyUse)',
|
||||
callback: async() =>{
|
||||
await cleanup()
|
||||
}
|
||||
},
|
||||
// Only show the reboot option if the server is running on a local network 仅在本地或局域网环境可重启服务
|
||||
isLocalNetwork(window.location.host) ? {
|
||||
content: rebootIcon.replace('currentColor','var(--error-color)') + ' '+ $t('Reboot ComfyUI') + ' (EasyUse)',
|
||||
callback: _ =>{
|
||||
if (confirm($t("Are you sure you'd like to reboot the server?"))){
|
||||
try {
|
||||
api.fetchApi("/easyuse/reboot");
|
||||
} catch (exception) {}
|
||||
}
|
||||
}
|
||||
} : null,
|
||||
);
|
||||
return options;
|
||||
};
|
||||
|
||||
let renderInfoEvent = LGraphCanvas.prototype.renderInfo
|
||||
if(disableRenderInfo){
|
||||
LGraphCanvas.prototype.renderInfo = function (ctx, x, y) {}
|
||||
}
|
||||
|
||||
if(!toolbar){
|
||||
toolbar = $el('div.easyuse-toolbar',[
|
||||
$el('div.easyuse-toolbar-item',{
|
||||
onclick:_=>{
|
||||
createGroupMap()
|
||||
}
|
||||
},[
|
||||
$el('div.easyuse-toolbar-icon.group', {innerHTML:groupIcon}),
|
||||
$el('div.easyuse-toolbar-tips',$t('Groups Map'))
|
||||
]),
|
||||
$el('div.easyuse-toolbar-item',{
|
||||
onclick:async()=>{
|
||||
await cleanup()
|
||||
}
|
||||
},[
|
||||
$el('div.easyuse-toolbar-icon.rocket',{innerHTML:rocketIcon}),
|
||||
$el('div.easyuse-toolbar-tips',$t('Cleanup Of GPU Usage'))
|
||||
]),
|
||||
])
|
||||
if(disableRenderInfo){
|
||||
toolbar.classList.add('disable-render-info')
|
||||
}else{
|
||||
toolbar.classList.remove('disable-render-info')
|
||||
}
|
||||
document.body.appendChild(toolbar)
|
||||
}
|
||||
|
||||
// rewrite handleFile
|
||||
let loadGraphDataEvent = app.loadGraphData
|
||||
app.loadGraphData = async function (data, clean=true) {
|
||||
// if(data?.extra?.cpr){
|
||||
// toast.copyright()
|
||||
// }
|
||||
if(data?.extra?.note){
|
||||
if(guideDialog) {
|
||||
guideDialog.remove()
|
||||
guideDialog = null
|
||||
}
|
||||
if(note && toolbar) toolbar.removeChild(note)
|
||||
const need_models = data.extra?.need_models || null
|
||||
guideDialog = new GuideDialog(data.extra.note, need_models)
|
||||
note = $el('div.easyuse-toolbar-item',{
|
||||
onclick:async()=>{
|
||||
guideDialog.toggle()
|
||||
}
|
||||
},[
|
||||
$el('div.easyuse-toolbar-icon.question',{innerHTML:quesitonIcon}),
|
||||
$el('div.easyuse-toolbar-tips',$t('Workflow Guide'))
|
||||
])
|
||||
if(toolbar) toolbar.insertBefore(note, toolbar.firstChild)
|
||||
}
|
||||
else{
|
||||
if(note) {
|
||||
toolbar.removeChild(note)
|
||||
note = null
|
||||
}
|
||||
}
|
||||
return await loadGraphDataEvent.apply(this, [...arguments])
|
||||
}
|
||||
|
||||
addToolBar(app)
|
||||
},
|
||||
beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if (nodeData.name.startsWith("easy")) {
|
||||
|
||||
@@ -0,0 +1,283 @@
|
||||
import { app } from "../../../../scripts/app.js";
|
||||
import { api } from "../../../../scripts/api.js";
|
||||
import { $el, ComfyDialog } from "../../../../scripts/ui.js";
|
||||
import { $t } from '../common/i18n.js'
|
||||
import { toast } from "../common/toast.js";
|
||||
import {sleep, accSub} from "../common/utils.js";
|
||||
|
||||
let api_keys = []
|
||||
let api_current = 0
|
||||
let user_info = {}
|
||||
|
||||
const api_cost = {
|
||||
'sd3': 6.5,
|
||||
'sd3-turbo': 4,
|
||||
}
|
||||
|
||||
class AccountDialog extends ComfyDialog {
|
||||
constructor() {
|
||||
super();
|
||||
this.lists = []
|
||||
this.dialog_div = null
|
||||
this.user_div = null
|
||||
}
|
||||
|
||||
addItem(index, user_div){
|
||||
return $el('div.easyuse-account-dialog-item',[
|
||||
$el('input',{type:'text',placeholder:'Enter name',oninput: e=>{
|
||||
const dataIndex = Array.prototype.indexOf.call(this.dialog_div.querySelectorAll('.easyuse-account-dialog-item'), e.target.parentNode)
|
||||
api_keys[dataIndex]['name'] = e.target.value
|
||||
},value:api_keys[index]['name']}),
|
||||
$el('input.key',{type:'text',oninput: e=>{
|
||||
const dataIndex = Array.prototype.indexOf.call(this.dialog_div.querySelectorAll('.easyuse-account-dialog-item'), e.target.parentNode)
|
||||
api_keys[dataIndex]['key'] = e.target.value
|
||||
},placeholder:'Enter APIKEY', value:api_keys[index]['key']}),
|
||||
$el('button.choose',{textContent:$t('Choose'),onclick:async(e)=>{
|
||||
const dataIndex = Array.prototype.indexOf.call(this.dialog_div.querySelectorAll('.easyuse-account-dialog-item'), e.target.parentNode)
|
||||
let name = api_keys[dataIndex]['name']
|
||||
let key = api_keys[dataIndex]['key']
|
||||
if(!name){
|
||||
toast.error($t('Please enter the account name'))
|
||||
return
|
||||
}
|
||||
else if(!key){
|
||||
toast.error($t('Please enter the APIKEY'))
|
||||
return
|
||||
}
|
||||
let missing = true
|
||||
for(let i=0;i<api_keys.length;i++){
|
||||
if(!api_keys[i].key) {
|
||||
missing = false
|
||||
break
|
||||
}
|
||||
}
|
||||
if(!missing){
|
||||
toast.error($t('APIKEY is not Empty'))
|
||||
return
|
||||
}
|
||||
// 保存记录
|
||||
api_current = dataIndex
|
||||
const body = new FormData();
|
||||
body.append('api_keys', JSON.stringify(api_keys));
|
||||
body.append('current',api_current)
|
||||
const res = await api.fetchApi('/easyuse/stability/set_api_keys', {
|
||||
method: 'POST',
|
||||
body
|
||||
})
|
||||
if (res.status == 200) {
|
||||
const data = await res.json()
|
||||
if(data?.account && data?.balance){
|
||||
const avatar = data.account?.profile_picture || null
|
||||
const email = data.account?.email || null
|
||||
const credits = data.balance?.credits || 0
|
||||
user_div.replaceChildren(
|
||||
$el('div.easyuse-account-user-info', {
|
||||
onclick:_=>{
|
||||
new AccountDialog().show(user_div);
|
||||
}
|
||||
},[
|
||||
$el('div.user',[
|
||||
$el('div.avatar', avatar ? [$el('img',{src:avatar})] : '😀'),
|
||||
$el('div.info', [
|
||||
$el('h5.name', email),
|
||||
$el('h6.remark','Credits: '+ credits)
|
||||
])
|
||||
]),
|
||||
$el('div.edit', {textContent:$t('Edit')})
|
||||
])
|
||||
)
|
||||
toast.success($t('Save Succeed'))
|
||||
}
|
||||
else toast.success($t('Save Succeed'))
|
||||
this.close()
|
||||
} else {
|
||||
toast.error($t('Save Failed'))
|
||||
}
|
||||
}}),
|
||||
$el('button.delete',{textContent:$t('Delete'),onclick:e=>{
|
||||
const dataIndex = Array.prototype.indexOf.call(this.dialog_div.querySelectorAll('.easyuse-account-dialog-item'), e.target.parentNode)
|
||||
if(api_keys.length<=1){
|
||||
toast.error($t('At least one account is required'))
|
||||
return
|
||||
}
|
||||
api_keys.splice(dataIndex,1)
|
||||
this.dialog_div.removeChild(e.target.parentNode)
|
||||
}}),
|
||||
])
|
||||
}
|
||||
|
||||
show(userdiv) {
|
||||
api_keys.forEach((item,index)=>{
|
||||
this.lists.push(this.addItem(index,userdiv))
|
||||
})
|
||||
this.dialog_div = $el("div.easyuse-account-dialog", this.lists)
|
||||
super.show(
|
||||
$el('div.easyuse-account-dialog-main',[
|
||||
$el('div',[
|
||||
$el('a',{href:'https://platform.stability.ai/account/keys',target:'_blank',textContent:$t('Getting Your APIKEY')}),
|
||||
]),
|
||||
this.dialog_div,
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
createButtons() {
|
||||
const btns = super.createButtons();
|
||||
btns.unshift($el('button',{
|
||||
type:'button',
|
||||
textContent:$t('Save Account Info'),
|
||||
onclick:_=>{
|
||||
let missing = true
|
||||
for(let i=0;i<api_keys.length;i++){
|
||||
if(!api_keys[i].key) {
|
||||
missing = false
|
||||
break
|
||||
}
|
||||
}
|
||||
if(!missing){
|
||||
toast.error($t('APIKEY is not Empty'))
|
||||
}
|
||||
else {
|
||||
const body = new FormData();
|
||||
body.append('api_keys', JSON.stringify(api_keys));
|
||||
api.fetchApi('/easyuse/stability/set_api_keys', {
|
||||
method: 'POST',
|
||||
body
|
||||
}).then(res => {
|
||||
if (res.status == 200) {
|
||||
toast.success($t('Save Succeed'))
|
||||
|
||||
} else {
|
||||
toast.error($t('Save Failed'))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}))
|
||||
btns.unshift($el('button',{
|
||||
type:'button',
|
||||
textContent:$t('Add Account'),
|
||||
onclick:_=>{
|
||||
const name = 'Account '+(api_keys.length).toString()
|
||||
api_keys.push({name,key:''})
|
||||
const item = this.addItem(api_keys.length - 1)
|
||||
this.lists.push(item)
|
||||
this.dialog_div.appendChild(item)
|
||||
}
|
||||
}))
|
||||
return btns
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
app.registerExtension({
|
||||
name: 'comfy.easyUse.account',
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if(nodeData.name == 'easy stableDiffusion3API'){
|
||||
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated;
|
||||
nodeType.prototype.onNodeCreated = async function() {
|
||||
onNodeCreated ? onNodeCreated?.apply(this, arguments) : undefined;
|
||||
const seed_widget = this.widgets.find(w => ['seed_num','seed'].includes(w.name))
|
||||
const seed_control = this.widgets.find(w=> ['control_before_generate','control_after_generate'].includes(w.name))
|
||||
let model_widget = this.widgets.find(w => w.name == 'model')
|
||||
model_widget.callback = value =>{
|
||||
cost_widget.value = '-'+api_cost[value]
|
||||
}
|
||||
const cost_widget = this.addWidget('text', 'cost_credit', '0', _=>{
|
||||
},{
|
||||
serialize:false,
|
||||
})
|
||||
cost_widget.disabled = true
|
||||
setTimeout(_=>{
|
||||
if(seed_control.name == 'control_before_generate' && seed_widget.value === 0){
|
||||
seed_widget.value = Math.floor(Math.random() * 4294967294)
|
||||
}
|
||||
cost_widget.value = '-'+api_cost[model_widget.value]
|
||||
},100)
|
||||
let user_div = $el('div.easyuse-account-user', [$t('Loading UserInfo...')])
|
||||
let account = this.addDOMWidget('account',"btn",$el('div.easyuse-account',user_div));
|
||||
// 更新balance信息
|
||||
api.addEventListener('stable-diffusion-api-generate-succeed', async ({detail}) => {
|
||||
let remarkDiv = user_div.querySelectorAll('.remark')
|
||||
if(remarkDiv && remarkDiv[0]){
|
||||
const credits = detail?.model ? api_cost[detail.model] : 0
|
||||
if(credits) {
|
||||
let balance = accSub(parseFloat(remarkDiv[0].innerText.replace(/Credits: /g,'')),credits)
|
||||
if(balance>0){
|
||||
remarkDiv[0].innerText = 'Credits: '+ balance.toString()
|
||||
}
|
||||
}
|
||||
}
|
||||
await sleep(10000)
|
||||
const res = await api.fetchApi('/easyuse/stability/balance')
|
||||
if(res.status == 200){
|
||||
const data = await res.json()
|
||||
if(data?.balance){
|
||||
const credits = data.balance?.credits || 0
|
||||
if(remarkDiv && remarkDiv[0]){
|
||||
remarkDiv[0].innerText = 'Credits: ' + credits
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
// 获取api_keys
|
||||
const res = await api.fetchApi('/easyuse/stability/api_keys')
|
||||
if (res.status == 200){
|
||||
let data = await res.json()
|
||||
api_keys = data.keys
|
||||
api_current = data.current
|
||||
if (api_keys.length > 0 && api_current!==undefined){
|
||||
const api_key = api_keys[api_current]['key']
|
||||
const api_name = api_keys[api_current]['name']
|
||||
if(!api_key){
|
||||
user_div.replaceChildren(
|
||||
$el('div.easyuse-account-user-info', {
|
||||
onclick:_=>{
|
||||
new AccountDialog().show(user_div);
|
||||
}
|
||||
},[
|
||||
$el('div.user',[
|
||||
$el('div.avatar', '😀'),
|
||||
$el('div.info', [
|
||||
$el('h5.name', api_name),
|
||||
$el('h6.remark',$t('Click to set the APIKEY first'))
|
||||
])
|
||||
]),
|
||||
$el('div.edit', {textContent:$t('Edit')})
|
||||
])
|
||||
)
|
||||
}else{
|
||||
// 获取账号信息
|
||||
const res = await api.fetchApi('/easyuse/stability/user_info')
|
||||
if(res.status == 200){
|
||||
const data = await res.json()
|
||||
if(data?.account && data?.balance){
|
||||
const avatar = data.account?.profile_picture || null
|
||||
const email = data.account?.email || null
|
||||
const credits = data.balance?.credits || 0
|
||||
user_div.replaceChildren(
|
||||
$el('div.easyuse-account-user-info', {
|
||||
onclick:_=>{
|
||||
new AccountDialog().show(user_div);
|
||||
}
|
||||
},[
|
||||
$el('div.user',[
|
||||
$el('div.avatar', avatar ? [$el('img',{src:avatar})] : '😀'),
|
||||
$el('div.info', [
|
||||
$el('h5.name', email),
|
||||
$el('h6.remark','Credits: '+ credits)
|
||||
])
|
||||
]),
|
||||
$el('div.edit', {textContent:$t('Edit')})
|
||||
])
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
+133
-116
@@ -1,116 +1,133 @@
|
||||
// import {app} from "/scripts/app.js";
|
||||
// import {api} from "/scripts/api.js";
|
||||
// import {$el} from "/scripts/ui.js";
|
||||
//
|
||||
// let script=document.createElement("script");
|
||||
// script.type="text/JavaScript";
|
||||
// script.innerHTML = `
|
||||
// function displayImage(e) {
|
||||
// console.log(e)
|
||||
// }
|
||||
// `
|
||||
// document.getElementsByTagName('head')[0].appendChild(script);
|
||||
//
|
||||
// const Loaders = ['easy fullLoader','easy a1111Loader','easy comfyLoader']
|
||||
// app.registerExtension({
|
||||
// name:"comfy.easyUse.contextMenu",
|
||||
// async setup(){
|
||||
// const existingContextMenu = LiteGraph.ContextMenu;
|
||||
// LiteGraph.ContextMenu = function(values,options){
|
||||
// const threshold = 20;
|
||||
// const enabled = true;
|
||||
// if(!enabled || (values?.length || 0) <= threshold || !(options?.callback) || values.some(i => typeof i !== 'string')){
|
||||
// if(enabled){
|
||||
// // console.log('Skipping context menu auto nesting for incompatible menu.');
|
||||
// }
|
||||
// return existingContextMenu.apply(this,[...arguments]);
|
||||
// }
|
||||
// const compatValues = values;
|
||||
// const originalValues = [...compatValues];
|
||||
// const folders = {};
|
||||
// const specialOps = [];
|
||||
// const folderless = [];
|
||||
// for(const value of compatValues){
|
||||
// const splitBy = value.indexOf('/') > -1 ? '/' : '\\';
|
||||
// const valueSplit = value.split(splitBy);
|
||||
// if(valueSplit.length > 1){
|
||||
// const key = valueSplit.shift();
|
||||
// folders[key] = folders[key] || [];
|
||||
// folders[key].push(valueSplit.join(splitBy));
|
||||
// }else if(value === 'CHOOSE' || value.startsWith('DISABLE ')){
|
||||
// specialOps.push(value);
|
||||
// }else{
|
||||
// folderless.push(value);
|
||||
// }
|
||||
// }
|
||||
// const foldersCount = Object.values(folders).length;
|
||||
// if(foldersCount > 0){
|
||||
// const oldcallback = options.callback;
|
||||
// options.callback = null;
|
||||
// const newCallback = (item,options) => {
|
||||
// if(['None','无','無','なし'].includes(item.content)) oldcallback('None',options)
|
||||
// else oldcallback(originalValues.find(i => i.endsWith(item.content),options));
|
||||
// };
|
||||
// const addContent = (content, folderName='') => {
|
||||
// const name = folderName ? `${folderName}/${content}` : content;
|
||||
// // 获取图像
|
||||
// // const imgRes = api.fetchApi(`/easyuse/model/thumbnail?name=${name}`)
|
||||
// // if (imgRes.status === 200) {
|
||||
// // let data = await imgRes.json();
|
||||
// // console.log(data)
|
||||
// // }
|
||||
//
|
||||
// const newContent = $el(
|
||||
// "span.easyuse-model",
|
||||
// {
|
||||
// $: (el) => {
|
||||
// el.onmousemove = (e) => {
|
||||
// console.log(1)
|
||||
// };
|
||||
// el.onmouseout = () => {
|
||||
// console.log(2)
|
||||
// // hiddenImage()
|
||||
// };
|
||||
// el.onmouseover = (e) => {
|
||||
// console.log(1)
|
||||
// // displayImage(el.dataset.imgName, styleName)
|
||||
// };
|
||||
// },
|
||||
// },content)
|
||||
//
|
||||
//
|
||||
// return {
|
||||
// content,
|
||||
// title:newContent.outerHTML,
|
||||
// callback: newCallback
|
||||
// }
|
||||
// }
|
||||
// const newValues = [];
|
||||
// for(const [folderName,folder] of Object.entries(folders)){
|
||||
// newValues.push({
|
||||
// content:folderName,
|
||||
// has_submenu:true,
|
||||
// callback:() => {},
|
||||
// submenu:{
|
||||
// options:folder.map(f => addContent(f,folderName)),
|
||||
// }
|
||||
// });
|
||||
// }
|
||||
// newValues.push(...folderless.map(f => ({
|
||||
// content:f,
|
||||
// callback:newCallback
|
||||
// })));
|
||||
// if(specialOps.length > 0)
|
||||
// newValues.push(...specialOps.map(f => ({
|
||||
// content:f,
|
||||
// callback:newCallback
|
||||
// })));
|
||||
// return existingContextMenu.call(this,newValues,options);
|
||||
// }
|
||||
// return existingContextMenu.apply(this,[...arguments]);
|
||||
// }
|
||||
// LiteGraph.ContextMenu.prototype = existingContextMenu.prototype;
|
||||
// },
|
||||
//
|
||||
// })
|
||||
//
|
||||
import {app} from "../../../../scripts/app.js";
|
||||
import {api} from "../../../../scripts/api.js";
|
||||
import {$el} from "../../../../scripts/ui.js";
|
||||
import {$t} from "../common/i18n.js";
|
||||
import {getExtension, spliceExtension} from '../common/utils.js'
|
||||
import {toast} from "../common/toast.js";
|
||||
|
||||
const setting_id = "Comfy.EasyUse.MenuNestSub"
|
||||
let enableMenuNestSub = false
|
||||
let thumbnails = []
|
||||
|
||||
export function addMenuNestSubSetting(app) {
|
||||
app.ui.settings.addSetting({
|
||||
id: setting_id,
|
||||
name: $t("Enable ContextMenu Auto Nest Subdirectories (ComfyUI-Easy-Use)"),
|
||||
type: "boolean",
|
||||
defaultValue: enableMenuNestSub,
|
||||
onChange(value) {
|
||||
enableMenuNestSub = !!value;
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
const getEnableMenuNestSub = _ => app.ui.settings.getSettingValue(setting_id, enableMenuNestSub)
|
||||
|
||||
|
||||
const Loaders = ['easy fullLoader','easy a1111Loader','easy comfyLoader']
|
||||
app.registerExtension({
|
||||
name:"comfy.easyUse.contextMenu",
|
||||
async setup(app){
|
||||
addMenuNestSubSetting(app)
|
||||
// 获取所有模型图像
|
||||
const imgRes = await api.fetchApi(`/easyuse/models/thumbnail`)
|
||||
if (imgRes.status === 200) {
|
||||
let data = await imgRes.json();
|
||||
thumbnails = data
|
||||
}
|
||||
else if(getEnableMenuNestSub()){
|
||||
toast.error($t("Too many thumbnails, have closed the display"))
|
||||
}
|
||||
const existingContextMenu = LiteGraph.ContextMenu;
|
||||
LiteGraph.ContextMenu = function(values,options){
|
||||
const threshold = 10;
|
||||
const enabled = getEnableMenuNestSub();
|
||||
if(!enabled || (values?.length || 0) <= threshold || !(options?.callback) || values.some(i => typeof i !== 'string')){
|
||||
if(enabled){
|
||||
// console.log('Skipping context menu auto nesting for incompatible menu.');
|
||||
}
|
||||
return existingContextMenu.apply(this,[...arguments]);
|
||||
}
|
||||
const compatValues = values;
|
||||
const originalValues = [...compatValues];
|
||||
const folders = {};
|
||||
const specialOps = [];
|
||||
const folderless = [];
|
||||
for(const value of compatValues){
|
||||
const splitBy = value.indexOf('/') > -1 ? '/' : '\\';
|
||||
const valueSplit = value.split(splitBy);
|
||||
if(valueSplit.length > 1){
|
||||
const key = valueSplit.shift();
|
||||
folders[key] = folders[key] || [];
|
||||
folders[key].push(valueSplit.join(splitBy));
|
||||
}else if(value === 'CHOOSE' || value.startsWith('DISABLE ')){
|
||||
specialOps.push(value);
|
||||
}else{
|
||||
folderless.push(value);
|
||||
}
|
||||
}
|
||||
const foldersCount = Object.values(folders).length;
|
||||
if(foldersCount > 0){
|
||||
const oldcallback = options.callback;
|
||||
options.callback = null;
|
||||
const newCallback = (item,options) => {
|
||||
if(['None','无','無','なし'].includes(item.content)) oldcallback('None',options)
|
||||
else oldcallback(originalValues.find(i => i.endsWith(item.content),options));
|
||||
};
|
||||
const addContent = (content, folderName='') => {
|
||||
const name = folderName ? folderName + '\\' + spliceExtension(content) : spliceExtension(content);
|
||||
const ext = getExtension(content)
|
||||
const time = new Date().getTime()
|
||||
let thumbnail = ''
|
||||
if(['ckpt', 'pt', 'bin', 'pth', 'safetensors'].includes(ext)){
|
||||
for(let i=0;i<thumbnails.length;i++){
|
||||
let thumb = thumbnails[i]
|
||||
if(name && thumb && thumb.indexOf(name) != -1){
|
||||
thumbnail = thumbnails[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let newContent
|
||||
if(thumbnail){
|
||||
const protocol = window.location.protocol
|
||||
const host = window.location.host
|
||||
const base_url = `${protocol}//${host}`
|
||||
const thumb_url = thumbnail.replace(':','%3A').replace(/\\/g,'/')
|
||||
newContent = $el("div.easyuse-model", {},[$el("span",{},content + ' *'),$el("img",{src:`${base_url}/${thumb_url}?t=${time}`})])
|
||||
}else{
|
||||
newContent = $el("div.easyuse-model", {},[
|
||||
$el("span",{},content)
|
||||
])
|
||||
}
|
||||
|
||||
return {
|
||||
content,
|
||||
title:newContent.outerHTML,
|
||||
callback: newCallback
|
||||
}
|
||||
}
|
||||
const newValues = [];
|
||||
for(const [folderName,folder] of Object.entries(folders)){
|
||||
newValues.push({
|
||||
content:folderName,
|
||||
has_submenu:true,
|
||||
callback:() => {},
|
||||
submenu:{
|
||||
options:folder.map(f => addContent(f,folderName)),
|
||||
}
|
||||
});
|
||||
}
|
||||
newValues.push(...folderless.map(f => addContent(f, '')));
|
||||
if(specialOps.length > 0)
|
||||
newValues.push(...specialOps.map(f => addContent(f, '')));
|
||||
return existingContextMenu.call(this,newValues,options);
|
||||
}
|
||||
return existingContextMenu.apply(this,[...arguments]);
|
||||
}
|
||||
LiteGraph.ContextMenu.prototype = existingContextMenu.prototype;
|
||||
},
|
||||
|
||||
})
|
||||
|
||||
|
||||
@@ -1,33 +1,14 @@
|
||||
import { app } from "/scripts/app.js";
|
||||
import { api } from "/scripts/api.js";
|
||||
import { ComfyWidgets } from "/scripts/widgets.js";
|
||||
import { app } from "../../../../scripts/app.js";
|
||||
import { api } from "../../../../scripts/api.js";
|
||||
import { ComfyWidgets } from "../../../../scripts/widgets.js";
|
||||
import { toast} from "../common/toast.js";
|
||||
import { $t } from '../common/i18n.js';
|
||||
|
||||
let origProps = {};
|
||||
import { findWidgetByName, toggleWidget, updateNodeHeight} from "../common/utils.js";
|
||||
|
||||
const findWidgetByName = (node, name) => node.widgets.find((w) => w.name === name);
|
||||
const seedNodes = ["easy seed", "easy latentNoisy", "easy wildcards", "easy preSampling", "easy preSamplingAdvanced", "easy preSamplingNoiseIn", "easy preSamplingSdTurbo", "easy preSamplingCascade", "easy preSamplingDynamicCFG", "easy preSamplingLayerDiffusion", "easy fullkSampler", "easy fullCascadeKSampler"]
|
||||
const loaderNodes = ["easy fullLoader", "easy a1111Loader", "easy comfyLoader"]
|
||||
|
||||
const doesInputWithNameExist = (node, name) => node.inputs ? node.inputs.some((input) => input.name === name) : false;
|
||||
|
||||
function updateNodeHeight(node) {
|
||||
node.setSize([node.size[0], node.computeSize()[1]]);
|
||||
}
|
||||
|
||||
function toggleWidget(node, widget, show = false, suffix = "") {
|
||||
if (!widget || doesInputWithNameExist(node, widget.name)) return;
|
||||
if (!origProps[widget.name]) {
|
||||
origProps[widget.name] = { origType: widget.type, origComputeSize: widget.computeSize };
|
||||
}
|
||||
const origSize = node.size;
|
||||
|
||||
widget.type = show ? origProps[widget.name].origType : "easyHidden" + suffix;
|
||||
widget.computeSize = show ? origProps[widget.name].origComputeSize : () => [0, -4];
|
||||
|
||||
widget.linkedWidgets?.forEach(w => toggleWidget(node, w, ":" + widget.name, show));
|
||||
|
||||
const height = show ? Math.max(node.computeSize()[1], origSize[1]) : node.size[1];
|
||||
node.setSize([node.size[0], height]);
|
||||
|
||||
}
|
||||
|
||||
function widgetLogic(node, widget) {
|
||||
if (widget.name === 'lora_name') {
|
||||
@@ -70,18 +51,18 @@ function widgetLogic(node, widget) {
|
||||
updateNodeHeight(node)
|
||||
}
|
||||
if (widget.name === 'image_output') {
|
||||
if (widget.value === 'Sender' || widget.value === 'Sender/Save'){
|
||||
if (widget.value === 'Sender' || widget.value === 'Sender&Save'){
|
||||
toggleWidget(node, findWidgetByName(node, 'link_id'), true)
|
||||
}else {
|
||||
toggleWidget(node, findWidgetByName(node, 'link_id'))
|
||||
}
|
||||
if (widget.value === 'Hide' || widget.value === 'Preview' || widget.value === 'Sender') {
|
||||
if (widget.value === 'Hide' || widget.value === 'Preview' || widget.value == 'Preview&Choose' || widget.value === 'Sender') {
|
||||
toggleWidget(node, findWidgetByName(node, 'save_prefix'))
|
||||
toggleWidget(node, findWidgetByName(node, 'output_path'))
|
||||
toggleWidget(node, findWidgetByName(node, 'embed_workflow'))
|
||||
toggleWidget(node, findWidgetByName(node, 'number_padding'))
|
||||
toggleWidget(node, findWidgetByName(node, 'overwrite_existing'))
|
||||
} else if (widget.value === 'Save' || widget.value === 'Hide/Save' || widget.value === 'Sender/Save') {
|
||||
} else if (widget.value === 'Save' || widget.value === 'Hide&Save' || widget.value === 'Sender&Save') {
|
||||
toggleWidget(node, findWidgetByName(node, 'save_prefix'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'output_path'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'embed_workflow'), true)
|
||||
@@ -89,7 +70,7 @@ function widgetLogic(node, widget) {
|
||||
toggleWidget(node, findWidgetByName(node, 'overwrite_existing'), true)
|
||||
}
|
||||
|
||||
if(widget.value === 'Hide' || widget.value === 'Hide/Save'){
|
||||
if(widget.value === 'Hide' || widget.value === 'Hide&Save'){
|
||||
toggleWidget(node, findWidgetByName(node, 'decode_vae_name'))
|
||||
}else{
|
||||
toggleWidget(node, findWidgetByName(node, 'decode_vae_name'), true)
|
||||
@@ -129,21 +110,72 @@ function widgetLogic(node, widget) {
|
||||
}
|
||||
updateNodeHeight(node)
|
||||
}
|
||||
if (widget.name === 'mode') {
|
||||
let number_to_show = findWidgetByName(node, 'num_loras').value + 1
|
||||
if (widget.name === 'num_controlnet') {
|
||||
let number_to_show = widget.value + 1
|
||||
for (let i = 0; i < number_to_show; i++) {
|
||||
if (widget.value === "simple") {
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_strength'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_model_strength'))
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_clip_strength'))
|
||||
toggleWidget(node, findWidgetByName(node, 'controlnet_'+i), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'controlnet_'+i+'_strength'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'scale_soft_weight_'+i),true)
|
||||
if (findWidgetByName(node, 'mode').value === "simple") {
|
||||
toggleWidget(node, findWidgetByName(node, 'start_percent_'+i))
|
||||
toggleWidget(node, findWidgetByName(node, 'end_percent_'+i))
|
||||
} else {
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_strength'))
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_model_strength'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_clip_strength'), true)}
|
||||
toggleWidget(node, findWidgetByName(node, 'start_percent_'+i),true)
|
||||
toggleWidget(node, findWidgetByName(node, 'end_percent_'+i), true)
|
||||
}
|
||||
}
|
||||
for (let i = number_to_show; i < 10; i++) {
|
||||
toggleWidget(node, findWidgetByName(node, 'controlnet_'+i))
|
||||
toggleWidget(node, findWidgetByName(node, 'controlnet_'+i+'_strength'))
|
||||
toggleWidget(node, findWidgetByName(node, 'start_percent_'+i))
|
||||
toggleWidget(node, findWidgetByName(node, 'end_percent_'+i))
|
||||
toggleWidget(node, findWidgetByName(node, 'scale_soft_weight_'+i))
|
||||
}
|
||||
updateNodeHeight(node)
|
||||
}
|
||||
|
||||
if (widget.name === 'mode') {
|
||||
switch (node.comfyClass) {
|
||||
case 'easy loraStack':
|
||||
for (let i = 0; i < (findWidgetByName(node, 'num_loras').value + 1); i++) {
|
||||
if (widget.value === "simple") {
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_strength'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_model_strength'))
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_clip_strength'))
|
||||
} else {
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_strength'))
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_model_strength'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_'+i+'_clip_strength'), true)}
|
||||
}
|
||||
updateNodeHeight(node)
|
||||
break
|
||||
case 'easy controlnetStack':
|
||||
for (let i = 0; i < (findWidgetByName(node, 'num_controlnet').value + 1); i++) {
|
||||
if (widget.value === "simple") {
|
||||
toggleWidget(node, findWidgetByName(node, 'start_percent_'+i))
|
||||
toggleWidget(node, findWidgetByName(node, 'end_percent_'+i))
|
||||
} else {
|
||||
toggleWidget(node, findWidgetByName(node, 'start_percent_' + i), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'end_percent_' + i), true)
|
||||
}
|
||||
}
|
||||
updateNodeHeight(node)
|
||||
break
|
||||
case 'easy icLightApply':
|
||||
if (widget.value === "Foreground") {
|
||||
toggleWidget(node, findWidgetByName(node, 'lighting'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'remove_bg'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'source'))
|
||||
} else {
|
||||
toggleWidget(node, findWidgetByName(node, 'lighting'))
|
||||
toggleWidget(node, findWidgetByName(node, 'source'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'remove_bg'))
|
||||
}
|
||||
updateNodeHeight(node)
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if (widget.name === 'resolution') {
|
||||
if (widget.value === "自定义 x 自定义") {
|
||||
toggleWidget(node, findWidgetByName(node, 'empty_latent_width'), true)
|
||||
@@ -152,7 +184,6 @@ function widgetLogic(node, widget) {
|
||||
toggleWidget(node, findWidgetByName(node, 'empty_latent_width'), false)
|
||||
toggleWidget(node, findWidgetByName(node, 'empty_latent_height'), false)
|
||||
}
|
||||
updateNodeHeight(node)
|
||||
}
|
||||
if (widget.name === 'downscale_mode') {
|
||||
const widget_names = ['block_number', 'downscale_factor', 'start_percent', 'end_percent', 'downscale_after_skip', 'downscale_method', 'upscale_method']
|
||||
@@ -208,6 +239,124 @@ function widgetLogic(node, widget) {
|
||||
toggleWidget(node, findWidgetByName(node, 'new_cond_end'), true)
|
||||
}
|
||||
}
|
||||
|
||||
if (widget.name === 'preset') {
|
||||
const normol_presets = [
|
||||
'LIGHT - SD1.5 only (low strength)',
|
||||
'STANDARD (medium strength)',
|
||||
'VIT-G (medium strength)',
|
||||
'PLUS (high strength)', 'PLUS FACE (portraits)',
|
||||
'FULL FACE - SD1.5 only (portraits stronger)',
|
||||
]
|
||||
const faceid_presets = [
|
||||
'FACEID',
|
||||
'FACEID PLUS - SD1.5 only',
|
||||
'FACEID PLUS V2',
|
||||
'FACEID PORTRAIT (style transfer)'
|
||||
]
|
||||
if(normol_presets.includes(widget.value)){
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_strength'))
|
||||
toggleWidget(node, findWidgetByName(node, 'provider'))
|
||||
toggleWidget(node, findWidgetByName(node, 'weight_faceidv2'))
|
||||
toggleWidget(node, findWidgetByName(node, 'use_tiled'), true)
|
||||
let use_tiled = findWidgetByName(node, 'use_tiled')
|
||||
if(use_tiled && use_tiled.value){
|
||||
toggleWidget(node, findWidgetByName(node, 'sharpening'), true)
|
||||
}else {
|
||||
toggleWidget(node, findWidgetByName(node, 'sharpening'))
|
||||
}
|
||||
|
||||
}
|
||||
else if(faceid_presets.includes(widget.value)){
|
||||
if(widget.value == 'FACEID PLUS V2'){
|
||||
toggleWidget(node, findWidgetByName(node, 'weight_faceidv2'), true)
|
||||
}else{
|
||||
toggleWidget(node, findWidgetByName(node, 'weight_faceidv2'))
|
||||
}
|
||||
if(widget.value == 'FACEID PORTRAIT (style transfer)'){
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_strength'), false)
|
||||
}
|
||||
else{
|
||||
toggleWidget(node, findWidgetByName(node, 'lora_strength'), true)
|
||||
}
|
||||
toggleWidget(node, findWidgetByName(node, 'provider'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'use_tiled'))
|
||||
toggleWidget(node, findWidgetByName(node, 'sharpening'))
|
||||
}
|
||||
updateNodeHeight(node)
|
||||
}
|
||||
|
||||
if (widget.name === 'use_tiled') {
|
||||
if(widget.value)
|
||||
toggleWidget(node, findWidgetByName(node, 'sharpening'), true)
|
||||
else
|
||||
toggleWidget(node, findWidgetByName(node, 'sharpening'))
|
||||
updateNodeHeight(node)
|
||||
}
|
||||
|
||||
if (widget.name === 'num_embeds') {
|
||||
let number_to_show = widget.value + 1
|
||||
for (let i = 0; i < number_to_show; i++) {
|
||||
toggleWidget(node, findWidgetByName(node, 'weight'+i), true)
|
||||
}
|
||||
for (let i = number_to_show; i < 6; i++) {
|
||||
toggleWidget(node, findWidgetByName(node, 'weight'+i))
|
||||
}
|
||||
updateNodeHeight(node)
|
||||
}
|
||||
|
||||
if (widget.name === 'guider'){
|
||||
switch (widget.value){
|
||||
case 'Basic':
|
||||
toggleWidget(node, findWidgetByName(node, 'cfg'))
|
||||
toggleWidget(node, findWidgetByName(node, 'cfg_negative'))
|
||||
break
|
||||
case 'CFG':
|
||||
toggleWidget(node, findWidgetByName(node, 'cfg'),true)
|
||||
toggleWidget(node, findWidgetByName(node, 'cfg_negative'))
|
||||
break
|
||||
case 'IP2P+DualCFG':
|
||||
case 'DualCFG':
|
||||
toggleWidget(node, findWidgetByName(node, 'cfg'),true)
|
||||
toggleWidget(node, findWidgetByName(node, 'cfg_negative'), true)
|
||||
break
|
||||
|
||||
}
|
||||
updateNodeHeight(node)
|
||||
}
|
||||
|
||||
if (widget.name === 'scheduler'){
|
||||
if (['karrasADV','exponentialADV','polyExponential'].includes(widget.value)){
|
||||
toggleWidget(node, findWidgetByName(node, 'sigma_max'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'sigma_min'), true)
|
||||
toggleWidget(node, findWidgetByName(node, 'denoise'))
|
||||
toggleWidget(node, findWidgetByName(node, 'beta_d'))
|
||||
toggleWidget(node, findWidgetByName(node, 'beta_min'))
|
||||
toggleWidget(node, findWidgetByName(node, 'eps_s'))
|
||||
if(widget.value != 'exponentialADV'){
|
||||
toggleWidget(node, findWidgetByName(node, 'rho'), true)
|
||||
}else{
|
||||
toggleWidget(node, findWidgetByName(node, 'rho'))
|
||||
}
|
||||
}else if(widget.value == 'vp'){
|
||||
toggleWidget(node, findWidgetByName(node, 'sigma_max'))
|
||||
toggleWidget(node, findWidgetByName(node, 'sigma_min'))
|
||||
toggleWidget(node, findWidgetByName(node, 'denoise'))
|
||||
toggleWidget(node, findWidgetByName(node, 'rho'))
|
||||
toggleWidget(node, findWidgetByName(node, 'beta_d'),true)
|
||||
toggleWidget(node, findWidgetByName(node, 'beta_min'),true)
|
||||
toggleWidget(node, findWidgetByName(node, 'eps_s'),true)
|
||||
}else{
|
||||
toggleWidget(node, findWidgetByName(node, 'denoise'),true)
|
||||
toggleWidget(node, findWidgetByName(node, 'sigma_max'))
|
||||
toggleWidget(node, findWidgetByName(node, 'sigma_min'))
|
||||
toggleWidget(node, findWidgetByName(node, 'beta_d'))
|
||||
toggleWidget(node, findWidgetByName(node, 'beta_min'))
|
||||
toggleWidget(node, findWidgetByName(node, 'eps_s'))
|
||||
toggleWidget(node, findWidgetByName(node, 'rho'))
|
||||
}
|
||||
updateNodeHeight(node)
|
||||
}
|
||||
}
|
||||
|
||||
function widgetLogic2(node, widget) {
|
||||
@@ -457,10 +606,14 @@ app.registerExtension({
|
||||
case "easy comfyLoader":
|
||||
case "easy cascadeLoader":
|
||||
case "easy svdLoader":
|
||||
case "easy dynamiCrafterLoader":
|
||||
case "easy loraStack":
|
||||
case "easy controlnetStack":
|
||||
case "easy latentNoisy":
|
||||
case "easy preSampling":
|
||||
case "easy preSamplingAdvanced":
|
||||
case "easy preSamplingNoiseIn":
|
||||
case "easy preSamplingCustom":
|
||||
case "easy preSamplingSdTurbo":
|
||||
case "easy preSamplingCascade":
|
||||
case "easy preSamplingLayerDiffusion":
|
||||
@@ -476,6 +629,9 @@ app.registerExtension({
|
||||
case "easy hiresFix":
|
||||
case "easy detailerFix":
|
||||
case "easy imageRemBg":
|
||||
case "easy imageColorMatch":
|
||||
case "easy imageDetailTransfer":
|
||||
case "easy loadImageBase64":
|
||||
case "easy XYInputs: Steps":
|
||||
case "easy XYInputs: Sampler/Scheduler":
|
||||
case 'easy XYInputs: Checkpoint':
|
||||
@@ -486,6 +642,10 @@ app.registerExtension({
|
||||
case "easy rangeFloat":
|
||||
case 'easy latentCompositeMaskedWithCond':
|
||||
case 'easy pipeEdit':
|
||||
case 'easy icLightApply':
|
||||
case 'easy ipadapterApply':
|
||||
case 'easy ipadapterApplyADV':
|
||||
case 'easy ipadapterApplyEncoder':
|
||||
getSetters(node)
|
||||
break
|
||||
case "easy wildcards":
|
||||
@@ -700,6 +860,7 @@ app.registerExtension({
|
||||
const pos = this.widgets.findIndex((w) => w.name === "spent_time");
|
||||
if (pos !== -1 && this.widgets[pos]) {
|
||||
const w = this.widgets[pos]
|
||||
console.log(text)
|
||||
w.value = text;
|
||||
}
|
||||
}
|
||||
@@ -734,7 +895,7 @@ app.registerExtension({
|
||||
};
|
||||
}
|
||||
|
||||
if (["easy fullLoader", "easy a1111Loader", "easy comfyLoader"].includes(nodeData.name)) {
|
||||
if (loaderNodes.includes(nodeData.name)) {
|
||||
function populate(text, type = 'positive') {
|
||||
if (this.widgets) {
|
||||
const pos = this.widgets.findIndex((w) => w.name === type + "_prompt");
|
||||
@@ -781,27 +942,57 @@ app.registerExtension({
|
||||
};
|
||||
}
|
||||
|
||||
if (["easy seed", "easy latentNoisy", "easy wildcards", "easy preSampling", "easy preSamplingAdvanced", "easy preSamplingNoiseIn", "easy preSamplingSdTurbo", "easy preSamplingCascade", "easy preSamplingDynamicCFG", "easy preSamplingLayerDiffusion", "easy fullkSampler", "easy fullCascadeKSampler"].includes(nodeData.name)) {
|
||||
if(["easy sv3dLoader"].includes(nodeData.name)){
|
||||
function changeSchedulerText(mode, batch_size, inputEl) {
|
||||
console.log(mode)
|
||||
switch (mode){
|
||||
case 'azimuth':
|
||||
inputEl.readOnly = true
|
||||
inputEl.style.opacity = 0.6
|
||||
return `0:(0.0,0.0)` + (batch_size > 1 ? `\n${batch_size-1}:(360.0,0.0)` : '')
|
||||
case 'elevation':
|
||||
inputEl.readOnly = true
|
||||
inputEl.style.opacity = 0.6
|
||||
return `0:(-90.0,0.0)` + (batch_size > 1 ? `\n${batch_size-1}:(90.0,0.0)` : '')
|
||||
case 'custom':
|
||||
inputEl.readOnly = false
|
||||
inputEl.style.opacity = 1
|
||||
return `0:(0.0,0.0)\n9:(180.0,0.0)\n20:(360.0,0.0)`
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated;
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
onNodeCreated ? onNodeCreated.apply(this, []) : undefined;
|
||||
const easing_mode_widget = this.widgets.find(w => w.name == 'easing_mode')
|
||||
const batch_size = this.widgets.find(w => w.name == 'batch_size')
|
||||
const scheduler = this.widgets.find(w => w.name == 'scheduler')
|
||||
setTimeout(_=>{
|
||||
if(!scheduler.value) scheduler.value = changeSchedulerText(easing_mode_widget.value, batch_size.value, scheduler.inputEl)
|
||||
},1)
|
||||
easing_mode_widget.callback = value=>{
|
||||
scheduler.value = changeSchedulerText(value, batch_size.value, scheduler.inputEl)
|
||||
}
|
||||
batch_size.callback = value =>{
|
||||
scheduler.value = changeSchedulerText(easing_mode_widget.value, value, scheduler.inputEl)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (seedNodes.includes(nodeData.name)) {
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated;
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
onNodeCreated ? onNodeCreated.apply(this, []) : undefined;
|
||||
// const values = ["randomize", "fixed", "increment", "decrement"]
|
||||
// const seed_widget = this.widgets.find(w => w.name == 'seed_num')
|
||||
// const seed_control = this.addWidget("combo", "control_before_generate", values[0], () => {
|
||||
// }, {
|
||||
// values,
|
||||
// serialize: false
|
||||
// })
|
||||
// seed_widget.linkedWidgets = [seed_control]
|
||||
const seed_widget = this.widgets.find(w => ['seed_num','seed'].includes(w.name))
|
||||
const seed_control = this.widgets.find(w=> ['control_before_generate','control_after_generate'].includes(w.name))
|
||||
if(nodeData.name == 'easy seed'){
|
||||
this.addWidget("button", "🎲 Manual Random Seed", null, _=>{
|
||||
if(seed_control.value != 'fixed'){
|
||||
seed_control.value = 'fixed'
|
||||
}
|
||||
const randomSeedButton = this.addWidget("button", "🎲 Manual Random Seed", null, _=>{
|
||||
if(seed_control.value != 'fixed') seed_control.value = 'fixed'
|
||||
seed_widget.value = Math.floor(Math.random() * 1125899906842624)
|
||||
})
|
||||
app.queuePrompt(0, 1)
|
||||
},{ serialize:false})
|
||||
seed_widget.linkedWidgets = [randomSeedButton, seed_control];
|
||||
}
|
||||
}
|
||||
const onAdded = nodeType.prototype.onAdded;
|
||||
@@ -809,9 +1000,11 @@ app.registerExtension({
|
||||
onAdded ? onAdded.apply(this, []) : undefined;
|
||||
const seed_widget = this.widgets.find(w => ['seed_num','seed'].includes(w.name))
|
||||
const seed_control = this.widgets.find(w=> ['control_before_generate','control_after_generate'].includes(w.name))
|
||||
if(seed_control.name == 'control_before_generate' && seed_widget.value === 0){
|
||||
seed_widget.value = Math.floor(Math.random() * 1125899906842624)
|
||||
}
|
||||
setTimeout(_=>{
|
||||
if(seed_control.name == 'control_before_generate' && seed_widget.value === 0) {
|
||||
seed_widget.value = Math.floor(Math.random() * 1125899906842624)
|
||||
}
|
||||
},1)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -842,7 +1035,7 @@ app.registerExtension({
|
||||
}
|
||||
}
|
||||
|
||||
if(['easy showAnything', 'easy showTensorShape'].includes(nodeData.name)){
|
||||
if(['easy showAnything', 'easy showTensorShape', 'easy imageInterrogator'].includes(nodeData.name)){
|
||||
function populate(text) {
|
||||
if (this.widgets) {
|
||||
const pos = this.widgets.findIndex((w) => w.name === "text");
|
||||
@@ -881,13 +1074,15 @@ app.registerExtension({
|
||||
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);
|
||||
}
|
||||
};
|
||||
if(!['easy imageInterrogator'].includes(nodeData.name)) {
|
||||
const onConfigure = nodeType.prototype.onConfigure;
|
||||
nodeType.prototype.onConfigure = function () {
|
||||
onConfigure?.apply(this, arguments);
|
||||
if (this.widgets_values?.length) {
|
||||
populate.call(this, this.widgets_values);
|
||||
}
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
if(nodeData.name == 'easy convertAnything'){
|
||||
@@ -913,6 +1108,36 @@ app.registerExtension({
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
if (nodeData.name == 'easy promptLine') {
|
||||
const onAdded = nodeType.prototype.onAdded;
|
||||
nodeType.prototype.onAdded = async function () {
|
||||
onAdded ? onAdded.apply(this, []) : undefined;
|
||||
let prompt_widget = this.widgets.find(w => w.name == "prompt")
|
||||
const button = this.addWidget("button", "get values from COMBO link", '', () => {
|
||||
const output_link = this.outputs[1]?.links?.length>0 ? this.outputs[1]['links'][0] : null
|
||||
const all_nodes = app.graph._nodes
|
||||
const node = all_nodes.find(cate=> cate.inputs?.find(input=> input.link == output_link))
|
||||
if(!output_link || !node){
|
||||
toast.error($t('No COMBO link'), 3000)
|
||||
return
|
||||
}
|
||||
else{
|
||||
const input = node.inputs.find(input=> input.link == output_link)
|
||||
const widget_name = input.widget.name
|
||||
const widgets = node.widgets
|
||||
const widget = widgets.find(cate=> cate.name == widget_name)
|
||||
let values = widget?.options.values || null
|
||||
if(values){
|
||||
values = values.join('\n')
|
||||
prompt_widget.value = values
|
||||
}
|
||||
}
|
||||
}, {
|
||||
serialize: false
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
@@ -922,7 +1147,11 @@ const getSetWidgets = ['rescale_after_model', 'rescale',
|
||||
'refiner_lora1_name', 'refiner_lora2_name', 'upscale_method',
|
||||
'image_output', 'add_noise', 'info', 'sampler_name',
|
||||
'ckpt_B_name', 'ckpt_C_name', 'save_model', 'refiner_ckpt_name',
|
||||
'num_loras', 'mode', 'toggle', 'resolution', 'target_parameter', 'input_count', 'replace_count', 'downscale_mode', 'range_mode','text_combine_mode', 'input_mode','lora_count','ckpt_count', 'conditioning_mode']
|
||||
'num_loras', 'num_controlnet', 'mode', 'toggle', 'resolution', 'target_parameter',
|
||||
'input_count', 'replace_count', 'downscale_mode', 'range_mode','text_combine_mode', 'input_mode',
|
||||
'lora_count','ckpt_count', 'conditioning_mode', 'preset', 'use_tiled', 'use_batch', 'num_embeds',
|
||||
"easing_mode", "guider", "scheduler"
|
||||
]
|
||||
|
||||
function getSetters(node) {
|
||||
if (node.widgets)
|
||||
|
||||
@@ -1,10 +1,15 @@
|
||||
import {app} from "/scripts/app.js";
|
||||
import {app} from "../../../../scripts/app.js";
|
||||
import {$t} from '../common/i18n.js'
|
||||
import {CheckpointInfoDialog, LoraInfoDialog} from "../common/model.js";
|
||||
|
||||
const loaders = ['easy fullLoader', 'easy a1111Loader', 'easy comfyLoader']
|
||||
const preSampling = ['easy preSampling', 'easy preSamplingAdvanced', 'easy preSamplingDynamicCFG', 'easy preSamplingNoiseIn', 'easy preSamplingLayerDiffusion', 'easy fullkSampler']
|
||||
const preSampling = ['easy preSampling', 'easy preSamplingAdvanced', 'easy preSamplingDynamicCFG', 'easy preSamplingNoiseIn', 'easy preSamplingCustom', 'easy preSamplingLayerDiffusion', 'easy fullkSampler']
|
||||
const kSampler = ['easy kSampler', 'easy kSamplerTiled', 'easy kSamplerInpainting', 'easy kSamplerDownscaleUnet', 'easy kSamplerLayerDiffusion']
|
||||
const controlnet = ['easy controlnetLoader', 'easy controlnetLoaderADV', 'easy instantIDApply', 'easy instantIDApplyADV']
|
||||
const ipadapter = ['easy ipadapterApply', 'easy ipadapterApplyADV', 'easy ipadapterStyleComposition', 'easy ipadapterApplyFromParams']
|
||||
const positive_prompt = ['easy positive', 'easy wildcards']
|
||||
const imageNode = ['easy loadImageBase64', 'LoadImage', 'LoadImageMask']
|
||||
const brushnet = ['easy applyBrushNet', 'easy applyPowerPaint']
|
||||
const widgetMapping = {
|
||||
"positive_prompt":{
|
||||
"text": "positive",
|
||||
@@ -45,6 +50,28 @@ const widgetMapping = {
|
||||
"cn_strength": ["strength", "cn_strength"],
|
||||
"cn_soft_weights": ["scale_soft_weights","cn_soft_weights"],
|
||||
},
|
||||
"ipadapter":{
|
||||
"preset":"preset",
|
||||
"lora_strength": "lora_strength",
|
||||
"provider": "provider",
|
||||
"weight":"weight",
|
||||
"weight_faceidv2": "weight_faceidv2",
|
||||
"start_at": "start_at",
|
||||
"end_at": "end_at",
|
||||
"cache_mode": "cache_mode",
|
||||
"use_tiled": "use_tiled",
|
||||
},
|
||||
"load_image":{
|
||||
"image":"image",
|
||||
"base64_data":"base64_data",
|
||||
"channel": "channel"
|
||||
},
|
||||
"brushnet":{
|
||||
"dtype": "dtype",
|
||||
"scale": "scale",
|
||||
"start_at": "start_at",
|
||||
"end_at": "end_at"
|
||||
}
|
||||
}
|
||||
const inputMapping = {
|
||||
"loaders":{
|
||||
@@ -73,6 +100,18 @@ const inputMapping = {
|
||||
"positive_prompt":{
|
||||
|
||||
},
|
||||
"ipadapter":{
|
||||
"model":"model",
|
||||
"image":"image",
|
||||
"image_style": "image",
|
||||
"attn_mask":"attn_mask",
|
||||
"optional_ipadapter":"optional_ipadapter"
|
||||
},
|
||||
"brushnet":{
|
||||
"pipe": "pipe",
|
||||
"image": "image",
|
||||
"mask": "mask"
|
||||
}
|
||||
};
|
||||
|
||||
const outputMapping = {
|
||||
@@ -104,6 +143,15 @@ const outputMapping = {
|
||||
"load_image":{
|
||||
"IMAGE":"IMAGE",
|
||||
"MASK": "MASK"
|
||||
},
|
||||
"ipadapter":{
|
||||
"model":"model",
|
||||
"tiles":"tiles",
|
||||
"masks":"masks",
|
||||
"ipadapter":"ipadapter"
|
||||
},
|
||||
"brushnet":{
|
||||
"pipe": "pipe",
|
||||
}
|
||||
};
|
||||
|
||||
@@ -240,6 +288,25 @@ const addMenu = (content, type, nodes_include, nodeType, has_submenu=true) => {
|
||||
has_submenu: has_submenu,
|
||||
callback: (value, options, e, menu, node) => showSwapMenu(value, options, e, menu, node, type, nodes_include)
|
||||
})
|
||||
if(type == 'loaders') {
|
||||
options.unshift({
|
||||
content: $t("💎 View Lora Info..."),
|
||||
callback: (value, options, e, menu, node) => {
|
||||
const widget = node.widgets.find(cate => cate.name == 'lora_name')
|
||||
let name = widget.value;
|
||||
if (!name || name == 'None') return
|
||||
new LoraInfoDialog(name).show('loras', name);
|
||||
}
|
||||
})
|
||||
options.unshift({
|
||||
content: $t("💎 View Checkpoint Info..."),
|
||||
callback: (value, options, e, menu, node) => {
|
||||
let name = node.widgets[0].value;
|
||||
if (!name || name == 'None') return
|
||||
new CheckpointInfoDialog(name).show('checkpoints', name);
|
||||
}
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
const showSwapMenu = (value, options, e, menu, node, type, nodes_include) => {
|
||||
@@ -458,7 +525,7 @@ app.registerExtension({
|
||||
// 刷新节点
|
||||
addMenuHandler(nodeType, function (_, options) {
|
||||
options.unshift({
|
||||
content: "🔃 Reload Node",
|
||||
content: $t("🔃 Reload Node"),
|
||||
callback: (value, options, e, menu, node) => {
|
||||
let graphcanvas = LGraphCanvas.active_canvas;
|
||||
if (!graphcanvas.selected_nodes || Object.keys(graphcanvas.selected_nodes).length <= 1) {
|
||||
@@ -470,7 +537,19 @@ app.registerExtension({
|
||||
}
|
||||
}
|
||||
})
|
||||
// ckptNames
|
||||
if(nodeData.name == 'easy ckptNames'){
|
||||
options.unshift({
|
||||
content: $t("💎 View Checkpoint Info..."),
|
||||
callback: (value, options, e, menu, node) => {
|
||||
let name = node.widgets[0].value;
|
||||
if (!name || name == 'None') return
|
||||
new CheckpointInfoDialog(name).show('checkpoints', name);
|
||||
}
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
// Swap提示词
|
||||
if (positive_prompt.includes(nodeData.name)) {
|
||||
addMenu("↪️ Swap EasyPrompt", 'positive_prompt', positive_prompt, nodeType)
|
||||
@@ -491,6 +570,18 @@ app.registerExtension({
|
||||
if (controlnet.includes(nodeData.name)) {
|
||||
addMenu("↪️ Swap EasyControlnet", 'controlnet', controlnet, nodeType)
|
||||
}
|
||||
// Swap IPAdapater
|
||||
if (ipadapter.includes(nodeData.name)) {
|
||||
addMenu("↪️ Swap EasyIPAdapater", 'ipadapter', ipadapter, nodeType)
|
||||
}
|
||||
// Swap Image
|
||||
if (imageNode.includes(nodeData.name)) {
|
||||
addMenu("↪️ Swap LoadImage", 'load_image', imageNode, nodeType)
|
||||
}
|
||||
// Swap Brushnet
|
||||
if (brushnet.includes(nodeData.name)) {
|
||||
addMenu("↪️ Swap BrushNet", 'brushnet', brushnet, nodeType)
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
|
||||
@@ -1,17 +1,18 @@
|
||||
import { app } from "/scripts/app.js";
|
||||
import { api } from "/scripts/api.js";
|
||||
import { $el } from "/scripts/ui.js";
|
||||
import { app } from "../../../../scripts/app.js";
|
||||
import { api } from "../../../../scripts/api.js";
|
||||
import { $el } from "../../../../scripts/ui.js";
|
||||
import {addPreconnect, addCss} from "../common/utils.js";
|
||||
|
||||
const locale = localStorage['AGL.Locale'] || localStorage['Comfy.Settings.AGL.Locale'] || 'en-US'
|
||||
|
||||
const customThemeColor = "#3f3eed"
|
||||
const customThemeColorLight = "#006691"
|
||||
const customThemeColorLight = "#008ecb"
|
||||
// 增加Slot颜色
|
||||
const customPipeLineLink = "#7737AA"
|
||||
const customPipeLineSDXLLink = "#7737AA"
|
||||
const customIntLink = "#29699C"
|
||||
const customXYPlotLink = "#74DA5D"
|
||||
const customLoraStackLink = "#94dccd"
|
||||
const customXYLink = "#38291f"
|
||||
|
||||
var customLinkColors = JSON.parse(localStorage.getItem('Comfy.Settings.ttN.customLinkColors')) || {};
|
||||
@@ -20,6 +21,8 @@ if (!customLinkColors["PIPE_LINE_SDXL"] || !LGraphCanvas.link_type_colors["PIPE_
|
||||
if (!customLinkColors["INT"] || !LGraphCanvas.link_type_colors["INT"]) {customLinkColors["INT"] = customIntLink;}
|
||||
if (!customLinkColors["XYPLOT"] || !LGraphCanvas.link_type_colors["XYPLOT"]) {customLinkColors["XYPLOT"] = customXYPlotLink;}
|
||||
if (!customLinkColors["X_Y"] || !LGraphCanvas.link_type_colors["X_Y"]) {customLinkColors["X_Y"] = customXYLink;}
|
||||
if (!customLinkColors["LORA_STACK"] || !LGraphCanvas.link_type_colors["LORA_STACK"]) {customLinkColors["LORA_STACK"] = customLoraStackLink;}
|
||||
if (!customLinkColors["CONTROL_NET_STACK"] || !LGraphCanvas.link_type_colors["CONTROL_NET_STACK"]) {customLinkColors["CONTROL_NET_STACK"] = customLoraStackLink;}
|
||||
|
||||
localStorage.setItem('Comfy.Settings.easyUse.customLinkColors', JSON.stringify(customLinkColors));
|
||||
|
||||
@@ -131,7 +134,6 @@ try{
|
||||
settings["AE.highlight"] = false
|
||||
}
|
||||
// 主题设置
|
||||
console.log(theme_name)
|
||||
if(!theme_name && _settings['Comfy.ColorPalette']) {
|
||||
theme_name = `"${_settings['Comfy.ColorPalette']}"`
|
||||
localStorage.setItem('Comfy.Settings.Comfy.ColorPalette', theme_name)
|
||||
@@ -743,8 +745,12 @@ const NODE_COLORS = {
|
||||
"easy positive":"green",
|
||||
"easy negative":"red",
|
||||
"easy promptList":"cyan",
|
||||
"easy promptLine":"cyan",
|
||||
"easy promptConcat":"cyan",
|
||||
"easy promptReplace":"cyan",
|
||||
"easy XYInputs: Seeds++ Batch": customXYLink,
|
||||
"easy XYInputs: ModelMergeBlocks": customXYLink,
|
||||
'easy textSwitch': "pale_blue"
|
||||
}
|
||||
|
||||
function setNodeColors(node, theme) {
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
// 1.0.2
|
||||
import { app } from "/scripts/app.js";
|
||||
import { GroupNodeConfig } from "/extensions/core/groupNode.js";
|
||||
import { api } from "/scripts/api.js";
|
||||
import { app } from "../../../../scripts/app.js";
|
||||
import { GroupNodeConfig } from "../../../../extensions/core/groupNode.js";
|
||||
import { api } from "../../../../scripts/api.js";
|
||||
import { $t } from "../common/i18n.js"
|
||||
|
||||
const nodeTemplateShortcutId = "Comfy.EasyUse.NodeTemplateShortcut"
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { app } from "/scripts/app.js";
|
||||
import { applyTextReplacements } from "/scripts/utils.js";
|
||||
import { app } from "../../../../scripts/app.js";
|
||||
import { applyTextReplacements } from "../../../../scripts/utils.js";
|
||||
|
||||
const extraNodes = ["easy imageSave", "easy fullkSampler", "easy kSampler", "easy kSamplerTiled","easy kSamplerInpainting", "easy kSamplerDownscaleUnet", "easy kSamplerSDTurbo","easy detailerFix"]
|
||||
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
import {app} from "../../../../scripts/app.js";
|
||||
import {$el} from "../../../../scripts/ui.js";
|
||||
import {$t} from "../common/i18n.js";
|
||||
import {findWidgetByName, toggleWidget} from "../common/utils.js";
|
||||
|
||||
|
||||
const tags = {
|
||||
"selfie_multiclass_256x256": ["Background", "Hair", "Body", "Face", "Clothes", "Others",],
|
||||
"human_parsing_lip":["Background","Hat","Hair","Glove","Sunglasses","Upper-clothes","Dress","Coat","Socks","Pants","Jumpsuits","Scarf","Skirt","Face","Left-arm","Right-arm","Left-leg","Right-leg","Left-shoe","Right-shoe"],
|
||||
}
|
||||
function getTagList(tags) {
|
||||
let rlist=[]
|
||||
tags.forEach((k,i) => {
|
||||
rlist.push($el(
|
||||
"label.easyuse-prompt-styles-tag",
|
||||
{
|
||||
dataset: {
|
||||
tag: i,
|
||||
name: $t(k),
|
||||
index: i
|
||||
},
|
||||
$: (el) => {
|
||||
el.children[0].onclick = () => {
|
||||
el.classList.toggle("easyuse-prompt-styles-tag-selected");
|
||||
};
|
||||
},
|
||||
},
|
||||
[
|
||||
$el("input",{
|
||||
type: 'checkbox',
|
||||
name: i
|
||||
}),
|
||||
$el("span",{
|
||||
textContent: $t(k),
|
||||
})
|
||||
]
|
||||
))
|
||||
});
|
||||
return rlist
|
||||
}
|
||||
|
||||
|
||||
app.registerExtension({
|
||||
name: 'comfy.easyUse.seg',
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
|
||||
if (nodeData.name == 'easy humanSegmentation') {
|
||||
// 创建时
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated;
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
onNodeCreated ? onNodeCreated?.apply(this, arguments) : undefined;
|
||||
const method = this.widgets.findIndex((w) => w.name == 'method');
|
||||
const list = $el("ul.easyuse-prompt-styles-list.no-top", []);
|
||||
let method_values = ''
|
||||
this.setProperty("values", [])
|
||||
|
||||
let selector = this.addDOMWidget('mask_components',"btn",$el('div.easyuse-prompt-styles',[list]))
|
||||
|
||||
Object.defineProperty(this.widgets[method],'value',{
|
||||
set:(value)=>{
|
||||
method_values = value
|
||||
if(method_values){
|
||||
selector.element.children[0].innerHTML = ''
|
||||
if(method_values == 'selfie_multiclass_256x256'){
|
||||
toggleWidget(this, findWidgetByName(this, 'confidence'), true)
|
||||
this.setSize([300, 260]);
|
||||
}else{
|
||||
toggleWidget(this, findWidgetByName(this, 'confidence'))
|
||||
this.setSize([300, 500]);
|
||||
}
|
||||
let list = getTagList(tags[method_values]);
|
||||
selector.element.children[0].append(...list)
|
||||
}
|
||||
},
|
||||
get: () => {
|
||||
return method_values
|
||||
}
|
||||
})
|
||||
|
||||
let mask_select_values = ''
|
||||
|
||||
Object.defineProperty(selector, "value", {
|
||||
set: (value) => {
|
||||
setTimeout(_=>{
|
||||
selector.element.children[0].querySelectorAll(".easyuse-prompt-styles-tag").forEach(el => {
|
||||
let arr = value.split(',')
|
||||
if (arr.includes(el.dataset.tag)) {
|
||||
el.classList.add("easyuse-prompt-styles-tag-selected");
|
||||
el.children[0].checked = true
|
||||
}
|
||||
})
|
||||
},100)
|
||||
},
|
||||
get: () => {
|
||||
selector.element.children[0].querySelectorAll(".easyuse-prompt-styles-tag").forEach(el => {
|
||||
if(el.classList.value.indexOf("easyuse-prompt-styles-tag-selected")>=0){
|
||||
if(!this.properties["values"].includes(el.dataset.tag)){
|
||||
this.properties["values"].push(el.dataset.tag);
|
||||
}
|
||||
}else{
|
||||
if(this.properties["values"].includes(el.dataset.tag)){
|
||||
this.properties["values"]= this.properties["values"].filter(v=>v!=el.dataset.tag);
|
||||
}
|
||||
}
|
||||
});
|
||||
mask_select_values = this.properties["values"].join(',');
|
||||
return mask_select_values;
|
||||
}
|
||||
});
|
||||
|
||||
let old_values = ''
|
||||
let mask_lists_dom = selector.element.children[0]
|
||||
|
||||
// 初始化
|
||||
setTimeout(_=>{
|
||||
if(!method_values) {
|
||||
method_values = 'selfie_multiclass_256x256'
|
||||
selector.element.children[0].innerHTML = ''
|
||||
// 重新排序
|
||||
let list = getTagList(tags[method_values]);
|
||||
selector.element.children[0].append(...list)
|
||||
}
|
||||
if(method_values == 'selfie_multiclass_256x256'){
|
||||
toggleWidget(this, findWidgetByName(this, 'confidence'), true)
|
||||
this.setSize([300, 260]);
|
||||
}else{
|
||||
toggleWidget(this, findWidgetByName(this, 'confidence'))
|
||||
this.setSize([300, 500]);
|
||||
}
|
||||
},1)
|
||||
|
||||
return onNodeCreated;
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
+17
-10
@@ -1,7 +1,8 @@
|
||||
// 1.0.3
|
||||
import { app } from "/scripts/app.js";
|
||||
import { api } from "/scripts/api.js";
|
||||
import { $el } from "/scripts/ui.js";
|
||||
import { app } from "../../../../scripts/app.js";
|
||||
import { api } from "../../../../scripts/api.js";
|
||||
import { $el } from "../../../../scripts/ui.js";
|
||||
import { $t } from "../common/i18n.js";
|
||||
|
||||
// 获取风格列表
|
||||
let styles_list_cache = {}
|
||||
@@ -91,10 +92,11 @@ async function displayImage(imgName, styleName) {
|
||||
img.src = empty_img
|
||||
}
|
||||
}
|
||||
var x = e.pageX-pxy.x-100;
|
||||
var y = e.pageY-pxy.y+25;
|
||||
img.style.left = x+"px";
|
||||
img.style.top = y+"px";
|
||||
var scale = app?.canvas?.ds?.scale || 1;
|
||||
var x = (e.pageX-pxy.x-100)/scale;
|
||||
var y = (e.pageY-pxy.y+25)/scale;
|
||||
img.style.left = x+"px";
|
||||
img.style.top = y+"px";
|
||||
img.style.display = "block";
|
||||
img.style.borderRadius = "10px";
|
||||
img.style.borderColor = "var(--fg-color)"
|
||||
@@ -118,20 +120,25 @@ app.registerExtension({
|
||||
onNodeCreated ? onNodeCreated?.apply(this, arguments) : undefined;
|
||||
const styles_id = this.widgets.findIndex((w) => w.name == 'styles');
|
||||
const language = localStorage['AGL.Locale'] || localStorage['Comfy.Settings.AGL.Locale'] || 'en-US'
|
||||
|
||||
const list = $el("ul.easyuse-prompt-styles-list",[]);
|
||||
let styles_values = ''
|
||||
this.setProperty("values", [])
|
||||
|
||||
let selector = this.addDOMWidget('select_styles',"btn",$el('div.easyuse-prompt-styles',[$el('div.tools', [
|
||||
$el('button.delete',{
|
||||
textContent: language == 'zh-CN' ? '清空所有' : 'Empty All',
|
||||
textContent: $t('Empty All'),
|
||||
style:{},
|
||||
onclick:()=>{
|
||||
selector.element.children[0].querySelectorAll(".search").forEach(el=>{
|
||||
el.value = ''
|
||||
})
|
||||
selector.element.children[1].querySelectorAll(".easyuse-prompt-styles-tag-selected").forEach(el => {
|
||||
el.classList.remove("easyuse-prompt-styles-tag-selected");
|
||||
el.children[0].checked = false
|
||||
})
|
||||
selector.element.children[1].querySelectorAll(".easyuse-prompt-styles-tag").forEach(el => {
|
||||
el.classList.remove('hide')
|
||||
})
|
||||
this.setProperty("values", [])
|
||||
}}
|
||||
),
|
||||
@@ -139,7 +146,7 @@ app.registerExtension({
|
||||
dir:"ltr",
|
||||
style:{"overflow-y": "scroll"},
|
||||
rows:1,
|
||||
placeholder:language == 'zh-CN' ? "🔎 在此处输入以搜索样式 ..." : "🔎 Type here to search styles ...",
|
||||
placeholder:$t("🔎 Type here to search styles ..."),
|
||||
oninput:(e)=>{
|
||||
let value = e.target.value
|
||||
selector.element.children[1].querySelectorAll(".easyuse-prompt-styles-tag").forEach(el => {
|
||||
|
||||
@@ -0,0 +1,523 @@
|
||||
import {app} from "../../../../scripts/app.js";
|
||||
import {api} from "../../../../scripts/api.js";
|
||||
import {$el} from "../../../../scripts/ui.js";
|
||||
|
||||
const propmts = ["easy wildcards", "easy positive", "easy negative", "easy stylesSelector", "easy promptConcat", "easy promptReplace"]
|
||||
const loaders = ["easy a1111Loader", "easy comfyLoader", "easy fullLoader", "easy svdLoader", "easy cascadeLoader", "easy sv3dLoader"]
|
||||
const preSamplingNodes = ["easy preSampling", "easy preSamplingAdvanced", "easy preSamplingNoiseIn", "easy preSamplingCustom", "easy preSamplingDynamicCFG","easy preSamplingSdTurbo", "easy preSamplingLayerDiffusion"]
|
||||
const kSampler = ["easy kSampler", "easy kSamplerTiled","easy kSamplerInpainting", "easy kSamplerDownscaleUnet", "easy kSamplerSDTurbo"]
|
||||
const controlNetNodes = ["easy controlnetLoader", "easy controlnetLoaderADV"]
|
||||
const instantIDNodes = ["easy instantIDApply", "easy instantIDApplyADV"]
|
||||
const ipadapterNodes = ["easy ipadapterApply", "easy ipadapterApplyADV" , "easy ipadapterStyleComposition"]
|
||||
const pipeNodes = ['easy pipeIn','easy pipeOut', 'easy pipeEdit']
|
||||
const xyNodes = ['easy XYPlot', 'easy XYPlotAdvanced']
|
||||
const extraNodes = ['easy setNode']
|
||||
const modelNormalNodes = [...["Reroute"],...['RescaleCFG','LoraLoaderModelOnly','LoraLoader','FreeU','FreeU_v2'],...ipadapterNodes,...extraNodes]
|
||||
const suggestions = {
|
||||
// prompt
|
||||
"easy seed":{
|
||||
"from":{
|
||||
"INT": [...["Reroute"],...preSamplingNodes,...['easy fullkSampler']]
|
||||
}
|
||||
},
|
||||
"easy positive":{
|
||||
"from":{
|
||||
"STRING": [...["Reroute"],...propmts]
|
||||
}
|
||||
},
|
||||
"easy negative":{
|
||||
"from":{
|
||||
"STRING": [...["Reroute"],...propmts]
|
||||
}
|
||||
},
|
||||
"easy wildcards":{
|
||||
"from":{
|
||||
"STRING": [...["Reroute","easy showAnything"],...propmts,]
|
||||
}
|
||||
},
|
||||
"easy stylesSelector":{
|
||||
"from":{
|
||||
"STRING": [...["Reroute","easy showAnything"],...propmts,]
|
||||
}
|
||||
},
|
||||
"easy promptConcat":{
|
||||
"from":{
|
||||
"STRING": [...["Reroute","easy showAnything"],...propmts,]
|
||||
}
|
||||
},
|
||||
"easy promptReplace":{
|
||||
"from":{
|
||||
"STRING": [...["Reroute","easy showAnything"],...propmts,]
|
||||
}
|
||||
},
|
||||
// sd相关
|
||||
"easy fullLoader": {
|
||||
"from":{
|
||||
"PIPE_LINE": [...["Reroute"],...preSamplingNodes,...['easy fullkSampler'],...pipeNodes,...extraNodes],
|
||||
"MODEL":modelNormalNodes
|
||||
},
|
||||
"to":{
|
||||
"STRING": [...["Reroute"],...propmts]
|
||||
}
|
||||
},
|
||||
"easy a1111Loader": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...preSamplingNodes, ...controlNetNodes, ...instantIDNodes, ...pipeNodes, ...extraNodes],
|
||||
"MODEL": modelNormalNodes
|
||||
},
|
||||
"to":{
|
||||
"STRING": [...["Reroute"],...propmts]
|
||||
}
|
||||
},
|
||||
"easy comfyLoader": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...preSamplingNodes, ...controlNetNodes, ...instantIDNodes, ...pipeNodes, ...extraNodes],
|
||||
"MODEL": modelNormalNodes
|
||||
},
|
||||
"to":{
|
||||
"STRING": [...["Reroute"],...propmts]
|
||||
}
|
||||
},
|
||||
"easy svdLoader":{
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...["easy preSampling", "easy preSamplingAdvanced", "easy preSamplingDynamicCFG"], ...pipeNodes, ...extraNodes],
|
||||
"MODEL": modelNormalNodes
|
||||
},
|
||||
"to":{
|
||||
"STRING": [...["Reroute"],...propmts]
|
||||
}
|
||||
},
|
||||
"easy zero123Loader":{
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...["easy preSampling", "easy preSamplingAdvanced", "easy preSamplingDynamicCFG"], ...pipeNodes, ...extraNodes],
|
||||
"MODEL": modelNormalNodes
|
||||
},
|
||||
"to":{
|
||||
"STRING": [...["Reroute"],...propmts]
|
||||
}
|
||||
},
|
||||
"easy sv3dLoader":{
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...["easy preSampling", "easy preSamplingAdvanced", "easy preSamplingDynamicCFG"], ...pipeNodes, ...extraNodes],
|
||||
"MODEL": modelNormalNodes
|
||||
},
|
||||
"to":{
|
||||
"STRING": [...["Reroute"],...propmts]
|
||||
}
|
||||
},
|
||||
"easy preSampling": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...kSampler, ...pipeNodes, ...controlNetNodes, ...xyNodes, ...extraNodes]
|
||||
},
|
||||
},
|
||||
"easy preSamplingAdvanced": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...kSampler, ...pipeNodes, ...controlNetNodes, ...xyNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
"easy preSamplingDynamicCFG": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...kSampler, ...pipeNodes, ...controlNetNodes, ...xyNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
"easy preSamplingCustom": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...kSampler, ...pipeNodes, ...controlNetNodes, ...xyNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
"easy preSamplingLayerDiffusion": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute", "easy kSamplerLayerDiffusion"], ...kSampler, ...pipeNodes, ...controlNetNodes, ...xyNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
"easy preSamplingNoiseIn": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...kSampler, ...pipeNodes, ...controlNetNodes, ...xyNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
// ksampler
|
||||
"easy fullkSampler": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...pipeNodes.reverse(), ...['easy preDetailerFix', 'easy preMaskDetailerFix'], ...preSamplingNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
"easy kSampler": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...pipeNodes.reverse(), ...['easy preDetailerFix', 'easy preMaskDetailerFix', 'easy hiresFix'], ...preSamplingNodes, ...extraNodes],
|
||||
}
|
||||
},
|
||||
// cn
|
||||
"easy controlnetLoader": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...preSamplingNodes, ...controlNetNodes, ...instantIDNodes, ...pipeNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
"easy controlnetLoaderADV":{
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...preSamplingNodes, ...controlNetNodes, ...instantIDNodes, ...pipeNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
// instant
|
||||
"easy instantIDApply": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...preSamplingNodes, ...controlNetNodes, ...instantIDNodes, ...pipeNodes, ...extraNodes],
|
||||
"MODEL": modelNormalNodes
|
||||
},
|
||||
"to":{
|
||||
"COMBO": [...["Reroute", "easy promptLine"]]
|
||||
}
|
||||
},
|
||||
"easy instantIDApplyADV":{
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...preSamplingNodes, ...controlNetNodes, ...instantIDNodes, ...pipeNodes, ...extraNodes],
|
||||
"MODEL": modelNormalNodes
|
||||
},
|
||||
"to":{
|
||||
"COMBO": [...["Reroute", "easy promptLine"]]
|
||||
}
|
||||
},
|
||||
"easy ipadapterApply":{
|
||||
"to":{
|
||||
"COMBO": [...["Reroute", "easy promptLine"]]
|
||||
}
|
||||
},
|
||||
"easy ipadapterApplyADV":{
|
||||
"to":{
|
||||
"COMBO": [...["Reroute", "easy promptLine"]]
|
||||
}
|
||||
},
|
||||
"easy ipadapterStyleComposition":{
|
||||
"to":{
|
||||
"COMBO": [...["Reroute", "easy promptLine"]]
|
||||
}
|
||||
},
|
||||
// fix
|
||||
"easy preDetailerFix":{
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute", "easy detailerFix"], ...pipeNodes, ...extraNodes]
|
||||
},
|
||||
"to":{
|
||||
"PIPE_LINE": [...["Reroute", "easy ultralyticsDetectorPipe", "easy samLoaderPipe", "easy kSampler", "easy fullkSampler"]]
|
||||
}
|
||||
},
|
||||
"easy preMaskDetailerFix":{
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute", "easy detailerFix"], ...pipeNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
"easy samLoaderPipe": {
|
||||
"from":{
|
||||
"PIPE_LINE": [...["Reroute", "easy preDetailerFix"], ...pipeNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
"easy ultralyticsDetectorPipe": {
|
||||
"from":{
|
||||
"PIPE_LINE": [...["Reroute", "easy preDetailerFix"], ...pipeNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
// cascade相关
|
||||
"easy cascadeLoader":{
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...["easy fullCascadeKSampler", 'easy preSamplingCascade'], ...controlNetNodes, ...pipeNodes, ...extraNodes],
|
||||
"MODEL": modelNormalNodes.filter(cate => !ipadapterNodes.includes(cate))
|
||||
}
|
||||
},
|
||||
"easy fullCascadeKSampler":{
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...["easy preSampling", "easy preSamplingAdvanced"], ...pipeNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
"easy preSamplingCascade":{
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...["easy cascadeKSampler",], ...pipeNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
"easy cascadeKSampler": {
|
||||
"from": {
|
||||
"PIPE_LINE": [...["Reroute"], ...["easy preSampling", "easy preSamplingAdvanced"], ...pipeNodes, ...extraNodes]
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
app.registerExtension({
|
||||
name: "comfy.easyuse.suggestions",
|
||||
async setup(app) {
|
||||
LGraphCanvas.prototype.createDefaultNodeForSlot = function(optPass) { // addNodeMenu for connection
|
||||
var optPass = optPass || {};
|
||||
var opts = Object.assign({ nodeFrom: null // input
|
||||
,slotFrom: null // input
|
||||
,nodeTo: null // output
|
||||
,slotTo: null // output
|
||||
,position: [] // pass the event coords
|
||||
,nodeType: null // choose a nodetype to add, AUTO to set at first good
|
||||
,posAdd:[0,0] // adjust x,y
|
||||
,posSizeFix:[0,0] // alpha, adjust the position x,y based on the new node size w,h
|
||||
}
|
||||
,optPass
|
||||
);
|
||||
var that = this;
|
||||
|
||||
var isFrom = opts.nodeFrom && opts.slotFrom!==null;
|
||||
var isTo = !isFrom && opts.nodeTo && opts.slotTo!==null;
|
||||
|
||||
if (!isFrom && !isTo){
|
||||
console.warn("No data passed to createDefaultNodeForSlot "+opts.nodeFrom+" "+opts.slotFrom+" "+opts.nodeTo+" "+opts.slotTo);
|
||||
return false;
|
||||
}
|
||||
if (!opts.nodeType){
|
||||
console.warn("No type to createDefaultNodeForSlot");
|
||||
return false;
|
||||
}
|
||||
|
||||
var nodeX = isFrom ? opts.nodeFrom : opts.nodeTo;
|
||||
var slotX = isFrom ? opts.slotFrom : opts.slotTo;
|
||||
var nodeType = nodeX.type
|
||||
|
||||
var iSlotConn = false;
|
||||
switch (typeof slotX){
|
||||
case "string":
|
||||
iSlotConn = isFrom ? nodeX.findOutputSlot(slotX,false) : nodeX.findInputSlot(slotX,false);
|
||||
slotX = isFrom ? nodeX.outputs[slotX] : nodeX.inputs[slotX];
|
||||
break;
|
||||
case "object":
|
||||
// ok slotX
|
||||
iSlotConn = isFrom ? nodeX.findOutputSlot(slotX.name) : nodeX.findInputSlot(slotX.name);
|
||||
break;
|
||||
case "number":
|
||||
iSlotConn = slotX;
|
||||
slotX = isFrom ? nodeX.outputs[slotX] : nodeX.inputs[slotX];
|
||||
break;
|
||||
case "undefined":
|
||||
default:
|
||||
// bad ?
|
||||
//iSlotConn = 0;
|
||||
console.warn("Cant get slot information "+slotX);
|
||||
return false;
|
||||
}
|
||||
|
||||
if (slotX===false || iSlotConn===false){
|
||||
console.warn("createDefaultNodeForSlot bad slotX "+slotX+" "+iSlotConn);
|
||||
}
|
||||
|
||||
// check for defaults nodes for this slottype
|
||||
var fromSlotType = slotX.type==LiteGraph.EVENT?"_event_":slotX.type;
|
||||
var slotTypesDefault = isFrom ? LiteGraph.slot_types_default_out : LiteGraph.slot_types_default_in;
|
||||
if(slotTypesDefault && slotTypesDefault[fromSlotType]){
|
||||
if (slotX.link !== null) {
|
||||
// is connected
|
||||
}else{
|
||||
// is not not connected
|
||||
}
|
||||
let nodeNewType = false;
|
||||
const fromOrTo = isFrom ? 'from' : 'to'
|
||||
if(suggestions[nodeType] && suggestions[nodeType][fromOrTo] && suggestions[nodeType][fromOrTo][fromSlotType]?.length>0){
|
||||
for(var typeX in suggestions[nodeType][fromOrTo][fromSlotType]){
|
||||
if (opts.nodeType == suggestions[nodeType][fromOrTo][fromSlotType][typeX] || opts.nodeType == "AUTO") {
|
||||
nodeNewType = suggestions[nodeType][fromOrTo][fromSlotType][typeX];
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
else if(typeof slotTypesDefault[fromSlotType] == "object" || typeof slotTypesDefault[fromSlotType] == "array"){
|
||||
for(var typeX in slotTypesDefault[fromSlotType]){
|
||||
if (opts.nodeType == slotTypesDefault[fromSlotType][typeX] || opts.nodeType == "AUTO"){
|
||||
nodeNewType = slotTypesDefault[fromSlotType][typeX];
|
||||
// console.log("opts.nodeType == slotTypesDefault[fromSlotType][typeX] :: "+opts.nodeType);
|
||||
break; // --------
|
||||
}
|
||||
}
|
||||
}else{
|
||||
if (opts.nodeType == slotTypesDefault[fromSlotType] || opts.nodeType == "AUTO") nodeNewType = slotTypesDefault[fromSlotType];
|
||||
}
|
||||
if (nodeNewType) {
|
||||
var nodeNewOpts = false;
|
||||
if (typeof nodeNewType == "object" && nodeNewType.node){
|
||||
nodeNewOpts = nodeNewType;
|
||||
nodeNewType = nodeNewType.node;
|
||||
}
|
||||
|
||||
//that.graph.beforeChange();
|
||||
|
||||
var newNode = LiteGraph.createNode(nodeNewType);
|
||||
if(newNode){
|
||||
// if is object pass options
|
||||
if (nodeNewOpts){
|
||||
if (nodeNewOpts.properties) {
|
||||
for (var i in nodeNewOpts.properties) {
|
||||
newNode.addProperty( i, nodeNewOpts.properties[i] );
|
||||
}
|
||||
}
|
||||
if (nodeNewOpts.inputs) {
|
||||
newNode.inputs = [];
|
||||
for (var i in nodeNewOpts.inputs) {
|
||||
newNode.addOutput(
|
||||
nodeNewOpts.inputs[i][0],
|
||||
nodeNewOpts.inputs[i][1]
|
||||
);
|
||||
}
|
||||
}
|
||||
if (nodeNewOpts.outputs) {
|
||||
newNode.outputs = [];
|
||||
for (var i in nodeNewOpts.outputs) {
|
||||
newNode.addOutput(
|
||||
nodeNewOpts.outputs[i][0],
|
||||
nodeNewOpts.outputs[i][1]
|
||||
);
|
||||
}
|
||||
}
|
||||
if (nodeNewOpts.title) {
|
||||
newNode.title = nodeNewOpts.title;
|
||||
}
|
||||
if (nodeNewOpts.json) {
|
||||
newNode.configure(nodeNewOpts.json);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
// add the node
|
||||
that.graph.add(newNode);
|
||||
newNode.pos = [ opts.position[0]+opts.posAdd[0]+(opts.posSizeFix[0]?opts.posSizeFix[0]*newNode.size[0]:0)
|
||||
,opts.position[1]+opts.posAdd[1]+(opts.posSizeFix[1]?opts.posSizeFix[1]*newNode.size[1]:0)]; //that.last_click_position; //[e.canvasX+30, e.canvasX+5];*/
|
||||
|
||||
//that.graph.afterChange();
|
||||
|
||||
// connect the two!
|
||||
if (isFrom){
|
||||
opts.nodeFrom.connectByType( iSlotConn, newNode, fromSlotType );
|
||||
}else{
|
||||
opts.nodeTo.connectByTypeOutput( iSlotConn, newNode, fromSlotType );
|
||||
}
|
||||
|
||||
// if connecting in between
|
||||
if (isFrom && isTo){
|
||||
// TODO
|
||||
}
|
||||
|
||||
return true;
|
||||
|
||||
}else{
|
||||
console.log("failed creating "+nodeNewType);
|
||||
}
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
LGraphCanvas.prototype.showConnectionMenu = function(optPass) { // addNodeMenu for connection
|
||||
var optPass = optPass || {};
|
||||
var opts = Object.assign({ nodeFrom: null // input
|
||||
,slotFrom: null // input
|
||||
,nodeTo: null // output
|
||||
,slotTo: null // output
|
||||
,e: null
|
||||
}
|
||||
,optPass
|
||||
);
|
||||
var that = this;
|
||||
|
||||
var isFrom = opts.nodeFrom && opts.slotFrom;
|
||||
var isTo = !isFrom && opts.nodeTo && opts.slotTo;
|
||||
|
||||
if (!isFrom && !isTo){
|
||||
console.warn("No data passed to showConnectionMenu");
|
||||
return false;
|
||||
}
|
||||
|
||||
var nodeX = isFrom ? opts.nodeFrom : opts.nodeTo;
|
||||
var slotX = isFrom ? opts.slotFrom : opts.slotTo;
|
||||
|
||||
var iSlotConn = false;
|
||||
switch (typeof slotX){
|
||||
case "string":
|
||||
iSlotConn = isFrom ? nodeX.findOutputSlot(slotX,false) : nodeX.findInputSlot(slotX,false);
|
||||
slotX = isFrom ? nodeX.outputs[slotX] : nodeX.inputs[slotX];
|
||||
break;
|
||||
case "object":
|
||||
// ok slotX
|
||||
iSlotConn = isFrom ? nodeX.findOutputSlot(slotX.name) : nodeX.findInputSlot(slotX.name);
|
||||
break;
|
||||
case "number":
|
||||
iSlotConn = slotX;
|
||||
slotX = isFrom ? nodeX.outputs[slotX] : nodeX.inputs[slotX];
|
||||
break;
|
||||
default:
|
||||
// bad ?
|
||||
//iSlotConn = 0;
|
||||
console.warn("Cant get slot information "+slotX);
|
||||
return false;
|
||||
}
|
||||
|
||||
var options = ["Add Node",null];
|
||||
|
||||
if (that.allow_searchbox){
|
||||
options.push("Search");
|
||||
options.push(null);
|
||||
}
|
||||
|
||||
// get defaults nodes for this slottype
|
||||
var fromSlotType = slotX.type==LiteGraph.EVENT?"_event_":slotX.type;
|
||||
var slotTypesDefault = isFrom ? LiteGraph.slot_types_default_out : LiteGraph.slot_types_default_in;
|
||||
var nodeType = nodeX.type
|
||||
if(slotTypesDefault && slotTypesDefault[fromSlotType]){
|
||||
const fromOrTo = isFrom ? 'from' : 'to'
|
||||
if(suggestions[nodeType] && suggestions[nodeType][fromOrTo] && suggestions[nodeType][fromOrTo][fromSlotType]?.length>0){
|
||||
for(var typeX in suggestions[nodeType][fromOrTo][fromSlotType]){
|
||||
options.push(suggestions[nodeType][fromOrTo][fromSlotType][typeX]);
|
||||
}
|
||||
}
|
||||
else if(typeof slotTypesDefault[fromSlotType] == "object" || typeof slotTypesDefault[fromSlotType] == "array"){
|
||||
for(var typeX in slotTypesDefault[fromSlotType]){
|
||||
options.push(slotTypesDefault[fromSlotType][typeX]);
|
||||
}
|
||||
}else{
|
||||
options.push(slotTypesDefault[fromSlotType]);
|
||||
}
|
||||
}
|
||||
|
||||
// build menu
|
||||
var menu = new LiteGraph.ContextMenu(options, {
|
||||
event: opts.e,
|
||||
title: (slotX && slotX.name!="" ? (slotX.name + (fromSlotType?" | ":"")) : "")+(slotX && fromSlotType ? fromSlotType : ""),
|
||||
callback: inner_clicked
|
||||
});
|
||||
|
||||
// callback
|
||||
function inner_clicked(v,options,e) {
|
||||
//console.log("Process showConnectionMenu selection");
|
||||
switch (v) {
|
||||
case "Add Node":
|
||||
LGraphCanvas.onMenuAdd(null, null, e, menu, function(node){
|
||||
if (isFrom){
|
||||
opts.nodeFrom.connectByType( iSlotConn, node, fromSlotType );
|
||||
}else{
|
||||
opts.nodeTo.connectByTypeOutput( iSlotConn, node, fromSlotType );
|
||||
}
|
||||
});
|
||||
break;
|
||||
case "Search":
|
||||
if(isFrom){
|
||||
that.showSearchBox(e,{node_from: opts.nodeFrom, slot_from: slotX, type_filter_in: fromSlotType});
|
||||
}else{
|
||||
that.showSearchBox(e,{node_to: opts.nodeTo, slot_from: slotX, type_filter_out: fromSlotType});
|
||||
}
|
||||
break;
|
||||
default:
|
||||
// check for defaults nodes for this slottype
|
||||
var nodeCreated = that.createDefaultNodeForSlot(Object.assign(opts,{ position: [opts.e.canvasX, opts.e.canvasY]
|
||||
,nodeType: v
|
||||
}));
|
||||
if (nodeCreated){
|
||||
// new node created
|
||||
//console.log("node "+v+" created")
|
||||
}else{
|
||||
// failed or v is not in defaults
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
return false;
|
||||
};
|
||||
}
|
||||
})
|
||||
@@ -1,5 +1,5 @@
|
||||
import { app } from "/scripts/app.js";
|
||||
import { ComfyWidgets } from "/scripts/widgets.js";
|
||||
import { app } from "../../../../scripts/app.js";
|
||||
import { ComfyWidgets } from "../../../../scripts/widgets.js";
|
||||
|
||||
const KEY_CODES = { ENTER: 13, ESC: 27, ARROW_DOWN: 40, ARROW_UP: 38 };
|
||||
const WIDGET_GAP = -4;
|
||||
@@ -150,7 +150,6 @@ const cssCode = `
|
||||
border-radius: 7px;
|
||||
text-align: center;
|
||||
text-wrap: balance;
|
||||
text-transform: uppercase;
|
||||
}
|
||||
.hideInfo-dropdown {
|
||||
position: absolute;
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { app } from "/scripts/app.js";
|
||||
import { app } from "../../../../scripts/app.js";
|
||||
import {removeDropdown, createDropdown} from "../common/dropdown.js";
|
||||
|
||||
function generateNumList(dictionary) {
|
||||
|
||||
+1
-2
@@ -1,5 +1,4 @@
|
||||
import { app } from "../../../scripts/app.js";
|
||||
import { ComfyWidgets } from '/scripts/widgets.js'
|
||||
|
||||
// Node that allows you to tunnel connections for cleaner graphs
|
||||
|
||||
@@ -201,7 +200,7 @@ app.registerExtension({
|
||||
},
|
||||
{
|
||||
values: () => {
|
||||
const setterNodes = graph._nodes.filter((otherNode) => otherNode.type == 'easy setNode');
|
||||
const setterNodes = node.graph._nodes.filter((otherNode) => otherNode.type == 'easy setNode');
|
||||
return setterNodes.map((otherNode) => otherNode.widgets[0].value).sort();
|
||||
}
|
||||
}
|
||||
|
||||
+2
-2
@@ -5,7 +5,7 @@ app.registerExtension({
|
||||
name: "comfy.easyUse.imageWidgets",
|
||||
|
||||
nodeCreated(node) {
|
||||
if (["easy imageSize","easy imageSizeBySide","easy imageSizeByLongerSide","easy imageSizeShow", "easy imagePixelPerfect"].includes(node.comfyClass)) {
|
||||
if (["easy imageSize","easy imageSizeBySide","easy imageSizeByLongerSide","easy imageSizeShow", "easy imageRatio", "easy imagePixelPerfect"].includes(node.comfyClass)) {
|
||||
|
||||
const inputEl = document.createElement("textarea");
|
||||
inputEl.className = "comfy-multiline-input";
|
||||
@@ -29,7 +29,7 @@ app.registerExtension({
|
||||
},
|
||||
|
||||
beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if (["easy imageSize","easy imageSizeBySide","easy imageSizeByLongerSide", "easy imageSizeShow", "easy imagePixelPerfect"].includes(nodeData.name)) {
|
||||
if (["easy imageSize","easy imageSizeBySide","easy imageSizeByLongerSide", "easy imageSizeShow", "easy imageRatio", "easy imagePixelPerfect"].includes(nodeData.name)) {
|
||||
function populate(arr_text) {
|
||||
var text = '';
|
||||
for (let i = 0; i < arr_text.length; i++){
|
||||
|
||||
@@ -0,0 +1,237 @@
|
||||
import { app } from "../../../../scripts/app.js";
|
||||
import { api } from "../../../../scripts/api.js";
|
||||
import { ComfyDialog, $el } from "../../../../scripts/ui.js";
|
||||
|
||||
import { restart_from_here } from "./prompt.js";
|
||||
import { hud, FlowState } from "./state.js";
|
||||
import { send_cancel, send_message, send_onstart, skip_next_restart_message } from "./messaging.js";
|
||||
import { display_preview_images, additionalDrawBackground, click_is_in_image } from "./preview.js";
|
||||
import {$t} from "../common/i18n.js";
|
||||
|
||||
|
||||
class chooserImageDialog extends ComfyDialog {
|
||||
|
||||
constructor() {
|
||||
super();
|
||||
this.node = null
|
||||
this.select_index = []
|
||||
this.dialog_div = null
|
||||
}
|
||||
|
||||
show(image,node){
|
||||
this.select_index = []
|
||||
this.node = node
|
||||
|
||||
const images_div = image.map((img, index) => {
|
||||
const imgEl = $el('img', {
|
||||
src: img.src,
|
||||
onclick: _ => {
|
||||
if(this.select_index.includes(index)){
|
||||
this.select_index = this.select_index.filter(i => i !== index)
|
||||
imgEl.classList.remove('selected')
|
||||
} else {
|
||||
this.select_index.push(index)
|
||||
imgEl.classList.add('selected')
|
||||
}
|
||||
if (node.selected.has(index)) node.selected.delete(index);
|
||||
else node.selected.add(index);
|
||||
}
|
||||
})
|
||||
return imgEl
|
||||
})
|
||||
super.show($el('div.easyuse-chooser-dialog',[
|
||||
$el('h5.easyuse-chooser-dialog-title', $t('Choose images to continue')),
|
||||
$el('div.easyuse-chooser-dialog-images',images_div)
|
||||
]))
|
||||
}
|
||||
createButtons() {
|
||||
const btns = super.createButtons();
|
||||
btns[0].onclick = _ => {
|
||||
if (FlowState.running()) { send_cancel();}
|
||||
super.close()
|
||||
}
|
||||
btns.unshift($el('button', {
|
||||
type: 'button',
|
||||
textContent: $t('Choose Selected Images'),
|
||||
onclick: _ => {
|
||||
if (FlowState.paused()) {
|
||||
send_message(this.node.id, [...this.node.selected, -1, ...this.node.anti_selected]);
|
||||
}
|
||||
if (FlowState.idle()) {
|
||||
skip_next_restart_message();
|
||||
restart_from_here(this.node.id).then(() => { send_message(this.node.id, [...this.node.selected, -1, ...this.node.anti_selected]); });
|
||||
}
|
||||
super.close()
|
||||
}
|
||||
}))
|
||||
return btns
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
function progressButtonPressed() {
|
||||
const node = app.graph._nodes_by_id[this.node_id];
|
||||
if (node) {
|
||||
const selected = [...node.selected]
|
||||
if(selected?.length>0){
|
||||
node.setProperty('values',selected)
|
||||
}
|
||||
if (FlowState.paused()) {
|
||||
send_message(node.id, [...node.selected, -1, ...node.anti_selected]);
|
||||
}
|
||||
if (FlowState.idle()) {
|
||||
skip_next_restart_message();
|
||||
restart_from_here(node.id).then(() => { send_message(node.id, [...node.selected, -1, ...node.anti_selected]); });
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function cancelButtonPressed() {
|
||||
|
||||
if (FlowState.running()) { send_cancel();}
|
||||
}
|
||||
|
||||
function enable_disabling(button) {
|
||||
Object.defineProperty(button, 'clicked', {
|
||||
get : function() { return this._clicked; },
|
||||
set : function(v) { this._clicked = (v && this.name!=''); }
|
||||
})
|
||||
}
|
||||
|
||||
function disable_serialize(widget) {
|
||||
if (!widget.options) widget.options = { };
|
||||
widget.options.serialize = false;
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name:'comfy.easyuse.imageChooser',
|
||||
init() {
|
||||
window.addEventListener("beforeunload", send_cancel, true);
|
||||
},
|
||||
setup(app) {
|
||||
|
||||
const draw = LGraphCanvas.prototype.draw;
|
||||
LGraphCanvas.prototype.draw = function() {
|
||||
if (hud.update()) {
|
||||
app.graph._nodes.forEach((node)=> { if (node.update) { node.update(); } })
|
||||
}
|
||||
draw.apply(this,arguments);
|
||||
}
|
||||
|
||||
|
||||
function easyuseImageChooser(event) {
|
||||
const {node,image,isKSampler} = display_preview_images(event);
|
||||
if(isKSampler) {
|
||||
const dialog = new chooserImageDialog();
|
||||
dialog.show(image,node)
|
||||
}
|
||||
}
|
||||
api.addEventListener("easyuse-image-choose", easyuseImageChooser);
|
||||
|
||||
/*
|
||||
If a run is interrupted, send a cancel message (unless we're doing the cancelling, to avoid infinite loop)
|
||||
*/
|
||||
const original_api_interrupt = api.interrupt;
|
||||
api.interrupt = function () {
|
||||
if (FlowState.paused() && !FlowState.cancelling) send_cancel();
|
||||
original_api_interrupt.apply(this, arguments);
|
||||
}
|
||||
|
||||
/*
|
||||
At the start of execution
|
||||
*/
|
||||
function on_execution_start() {
|
||||
if (send_onstart()) {
|
||||
app.graph._nodes.forEach((node)=> {
|
||||
if (node.selected || node.anti_selected) {
|
||||
node.selected.clear();
|
||||
node.anti_selected.clear();
|
||||
node.update();
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
api.addEventListener("execution_start", on_execution_start);
|
||||
},
|
||||
|
||||
async nodeCreated(node, app) {
|
||||
|
||||
if(node.comfyClass == 'easy imageChooser'){
|
||||
node.setProperty('values',[])
|
||||
|
||||
/* A property defining the top of the image when there is just one */
|
||||
if(node?.imageIndex === undefined){
|
||||
Object.defineProperty(node, 'imageIndex', {
|
||||
get : function() { return null; },
|
||||
set: function (v) {node.overIndex= v},
|
||||
})
|
||||
}
|
||||
if(node?.imagey === undefined){
|
||||
Object.defineProperty(node, 'imagey', {
|
||||
get : function() { return null; },
|
||||
set: function (v) {return node.widgets[node.widgets.length-1].last_y+LiteGraph.NODE_WIDGET_HEIGHT;},
|
||||
})
|
||||
}
|
||||
|
||||
/* Capture clicks */
|
||||
const org_onMouseDown = node.onMouseDown;
|
||||
node.onMouseDown = function( e, pos, canvas ) {
|
||||
if (e.isPrimary) {
|
||||
const i = click_is_in_image(node, pos);
|
||||
if (i>=0) { this.imageClicked(i); }
|
||||
}
|
||||
return (org_onMouseDown && org_onMouseDown.apply(this, arguments));
|
||||
}
|
||||
|
||||
node.send_button_widget = node.addWidget("button", "", "", progressButtonPressed);
|
||||
node.cancel_button_widget = node.addWidget("button", "", "", cancelButtonPressed);
|
||||
enable_disabling(node.cancel_button_widget);
|
||||
enable_disabling(node.send_button_widget);
|
||||
disable_serialize(node.cancel_button_widget);
|
||||
disable_serialize(node.send_button_widget);
|
||||
|
||||
}
|
||||
},
|
||||
|
||||
beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if(nodeData?.name == 'easy imageChooser'){
|
||||
|
||||
const onDrawBackground = nodeType.prototype.onDrawBackground;
|
||||
nodeType.prototype.onDrawBackground = function(ctx) {
|
||||
onDrawBackground.apply(this, arguments);
|
||||
additionalDrawBackground(this, ctx);
|
||||
}
|
||||
|
||||
nodeType.prototype.imageClicked = function (imageIndex) {
|
||||
if (nodeType?.comfyClass==="easy imageChooser") {
|
||||
if (this.selected.has(imageIndex)) this.selected.delete(imageIndex);
|
||||
else this.selected.add(imageIndex);
|
||||
this.update();
|
||||
}
|
||||
}
|
||||
|
||||
const update = nodeType.prototype.update;
|
||||
nodeType.prototype.update = function() {
|
||||
if (update) update.apply(this,arguments);
|
||||
if (this.send_button_widget) {
|
||||
this.send_button_widget.node_id = this.id;
|
||||
const selection = ( this.selected ? this.selected.size : 0 ) + ( this.anti_selected ? this.anti_selected.size : 0 )
|
||||
const maxlength = this.imgs?.length || 0;
|
||||
if (FlowState.paused_here(this.id) && selection>0) {
|
||||
this.send_button_widget.name = (selection>1) ? "Progress selected (" + selection + '/' + maxlength +")" : "Progress selected image";
|
||||
} else if (selection>0) {
|
||||
this.send_button_widget.name = (selection>1) ? "Progress selected (" + selection + '/' + maxlength +")" : "Progress selected image as restart";
|
||||
}
|
||||
else {
|
||||
this.send_button_widget.name = "";
|
||||
}
|
||||
}
|
||||
if (this.cancel_button_widget) {
|
||||
const isRunning = FlowState.running()
|
||||
this.cancel_button_widget.name = isRunning ? "Cancel current run" : "";
|
||||
}
|
||||
this.setDirtyCanvas(true,true);
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user