Compare commits

...
96 Commits
Author SHA1 Message Date
yolain 38851372e1 Upgrade version 1.1.8 to comfyregistry 2024-06-04 23:32:31 +08:00
yolain 1e9ffc5ffc add:auto translate chinese prompt to english 2024-06-03 14:44:28 +08:00
yolain ccb17f18b8 fix:xyplot error #192 2024-06-03 10:19:50 +08:00
yolain 380c596d9a fix:easy preSamplingCustom error 2024-06-02 02:21:49 +08:00
yolain 4c3328797b fix:compatibility powerpaint and brushnet new version 2024-06-01 15:35:55 +08:00
yolain d22f1f44f8 add:remove backend cache when clean gpu used 2024-05-31 15:51:03 +08:00
yolain fb435d47ea fix:easy imageChooser can not cancel queue 2024-05-29 19:08:58 +08:00
yolain d9d597bf83 fix:image object has not attribute movedim in layerDiffuse 2024-05-28 21:58:48 +08:00
yolain 1913c65d6f add:swapper for brushnet&powerpaint 2024-05-27 20:49:21 +08:00
yolain b8b24040eb add:easy controlnetStack 2024-05-26 21:40:01 +08:00
yolain e9bed88d63 optimized code for easy loader 2024-05-25 18:51:38 +08:00
yolain d95772147f 💖Upgrade to v1.1.8 2024-05-25 11:47:51 +08:00
yolain 7cbc2a2a2b modify the pyproject.toml file 2024-05-23 13:33:03 +08:00
yolain 63811206c9 Merge pull request #182 from haohaocreates/publish
Add Github Action for Publishing to Comfy Registry
2024-05-23 13:30:16 +08:00
yolain 48a4b5bfc6 Merge pull request #183 from haohaocreates/pyproject
Add pyproject.toml for Custom Node Registry
2024-05-23 13:30:00 +08:00
yolain 7980f4eef4 fix:compatible with new brushnet versions #181 2024-05-23 12:33:45 +08:00
haohaocreates b19c8b07ea chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-05-22 18:08:19 -04:00
haohaocreates 3f506d265c chore(publish): Add Github Action for Publishing to Comfy Registry 2024-05-22 18:08:16 -04:00
yolain d8af918e7d fix:the modal for selecting an image can not be closed #180 2024-05-21 15:54:02 +08:00
yolain 33566e8474 fix:run error when input optional image in easy fullkSampler #178 2024-05-21 15:41:39 +08:00
yolain ea0350c2dc fix:preview image offset in easy stylesSelector when zooming the page #176 2024-05-18 15:35:28 +08:00
yolain 24526623fb fix:iclight cache bug #173 2024-05-17 17:00:39 +08:00
yolain 7715ebfd06 Merge pull request #170 from chenpx976/main
feat: optimize GPU memory management and model unloading
2024-05-16 12:52:09 +08:00
color 2ba814c131 feat: optimize GPU memory management and model unloading 2024-05-16 12:35:40 +08:00
yolain d5ad332666 fix:icLightModel add to cache 2024-05-16 11:57:19 +08:00
yolain fee7e7bf73 fix:some models were not successfully written to easyCache,resulting in slow secondary diffusion 2024-05-16 01:01:03 +08:00
yolain 01f17ff02b remove:unnecessary import modules 2024-05-15 17:07:14 +08:00
yolain da7120219a fix:diffusers version>0.26.0 use different module #70 2024-05-15 10:50:19 +08:00
yolain d616d18069 fix:compatible with cg-image-picker #169 2024-05-15 10:04:00 +08:00
yolain dbf76f288c Update README.md 2024-05-14 16:49:16 +08:00
yolain 05124006ba fix:set_clip_options no longer compare version #165 2024-05-13 14:23:55 +08:00
yolain 4586af311c fix:fooocus+dd wrong 2024-05-13 14:19:33 +08:00
yolain 1cea58c7cf add:remove_bg in easy icLightApply 2024-05-11 02:13:55 +08:00
yolain a84f7c4a58 fix:brushnet error in easy presamplingInpainting 2024-05-11 00:42:14 +08:00
yolain 1d9bf86560 add:easy icLightApply 2024-05-10 17:26:04 +08:00
yolain 8d352b85bc add:easy imageSplitGrid 2024-05-09 16:08:03 +08:00
yolain 403562575c add:easy imageCropFromMask & imageUncropFromBBOX 2024-05-06 12:52:37 +08:00
yolain dbe2cd6569 fix:compatibility with comfyui-brushnet new commit 2024-05-05 11:30:41 +08:00
yolain b3b0a961c5 fix:raise excenption when comfyui-brushnet is not installed 2024-05-04 19:26:18 +08:00
yolain 9dfb8b9c15 fix:layerdiffuse everything config #158 2024-05-04 19:19:01 +08:00
yolain b4ea58946c fix:supported brushnet for ays 2024-05-04 01:01:35 +08:00
yolain 513fc4b67e support for brushnet model loading 2024-05-03 18:24:47 +08:00
yolain 924a16e31c fix:easy kSamplerInpainting and js file set full relative path 2024-05-02 12:01:56 +08:00
yolain ad43f8e0bb fix:preview&choose return new batch #153 2024-04-30 17:24:21 +08:00
yolain b7e1ce8a3c fix:easy imageChooser is not working #150 2024-04-29 12:03:09 +08:00
yolain 299090d184 add:denoise value to alignYourStepsScheduler #146 2024-04-29 11:32:49 +08:00
yolain 3aa7bca86b add:align_your_steps in all easy preSampling 2024-04-26 19:39:26 +08:00
yolain 56de6f0bc9 add:alignYourSteps of scheduler in easy preSamplingCustom #146 2024-04-26 15:35:06 +08:00
yolain f6b5f5c99c Upgrade to v1.1.6 2024-04-26 11:55:12 +08:00
yolain 79b81100ff fix:easy ipadapterApply some changes 2024-04-26 09:19:03 +08:00
yolain c885fd9bcf fix:easy ipadapterApply bug 2024-04-25 11:02:20 +08:00
yolain c46aaa6084 fix:incorrect string preview when refresh page 2024-04-25 00:19:37 +08:00
yolain 0562d4eb0e add:easy humanSegmentation 2024-04-24 22:47:33 +08:00
yolain cfa56d36d7 add:easy imageColorMatch 2024-04-24 17:37:48 +08:00
yolain 913cfe73ae rewrite:easy cleanGPUUsed can force cleanup of models gpu usage 2024-04-24 16:56:27 +08:00
yolain 590161c560 add:easy ipadadpterApplyFromParams 2024-04-23 14:39:38 +08:00
yolain d674313240 fix:api-key placeholder 2024-04-19 12:19:26 +08:00
yolain 44258090cd fix:deduct credit when request successful in easy stablediffusion3API 2024-04-19 11:53:49 +08:00
yolain 1222799e8f add:stableDiffusion3 API node 2024-04-19 11:20:15 +08:00
yolain ecc972f7b4 fix:easy preSamplingCustom optional_sampler and optional_sigmas 2024-04-18 12:20:51 +08:00
yolain c079652878 fix:ipadapterApplyADV use_tiled error #135 2024-04-16 13:03:51 +08:00
yolain b333bf05c6 fix:wildcard not match space #131 2024-04-15 19:26:46 +08:00
yolain f52a8fd40b add:easy preSamplingCustom 2024-04-13 11:59:20 +08:00
yolain f320647d78 fix:set the default value of contextmenu auto nest subdirectories to be disabled 2024-04-12 08:43:15 +08:00
yolain 631acfe44a fix:ipadapterApply error 2024-04-10 15:59:20 +08:00
yolain 7269f02f8a rename:clear cache key and clear cache all 2024-04-09 23:58:37 +08:00
yolain fdc761ebfa add:easy ipadapterStyleComposition 2024-04-09 23:48:41 +08:00
yolain 689d988130 fix:easy ipadapterApplyADV compatible with new version #123 2024-04-08 12:32:34 +08:00
yolain e29fd5ed24 Upgrade to v1.1.4 2024-04-07 12:43:09 +08:00
yolain 35f19b75fa fix:latentNoisy & preSamplingNoiseIn & Unsampler 2024-04-07 12:40:24 +08:00
yolain 96125a65a0 fix:getset.js 2024-04-06 10:29:49 +08:00
yolain 70ca7cc35f add:COMOSITION of preset in easy ipadapterApply 2024-04-05 12:33:18 +08:00
yolain 5858fb0606 add:supported load resadapter lora #116 2024-04-04 16:53:16 +08:00
yolain 149fab2105 add:if the huggingface connection timeout, it will switch to the mirrored address to download 2024-04-03 19:36:08 +08:00
yolain 39c5ccf469 rewrite:nodes suggestions in easyuse nodes 2024-04-03 18:01:07 +08:00
yolain cd5fcd1f70 add:enable contextmenu auto nest subdirectories 2024-04-03 14:59:23 +08:00
yolain 96b3897f95 fix:Differential Diffusion with grow_mask_by in kSamplerInpainting 2024-04-03 00:39:02 +08:00
yolain 293692398e fix:fooocus_model can not override the original incoming model #74 2024-04-03 00:12:36 +08:00
yolain 84307b95aa fix:lora not use different checkpoint model when optional_lora_stack in xyplot #113 2024-04-02 13:22:54 +08:00
yolain 1e1b82399c add:easy imageRatio #111 2024-04-01 19:19:16 +08:00
yolain 8c610dd8f9 fix:supported faceid portrait sdxl 2024-04-01 11:07:30 +08:00
yolain ac7db29df9 add:easy sv3dLoader 2024-03-31 18:19:16 +08:00
yolain d5fc43e11f fix:unable load batch_size value in easy pipeIn when link the image and pipe is None 2024-03-31 11:31:39 +08:00
yolain 65ba451d5d rename:dynamiCrafter dir 2024-03-30 11:44:03 +08:00
yolain cc79b11c12 fix:clear the search field when clicking on empty all in easy styleSelector node #109 2024-03-30 11:39:52 +08:00
yolain 148b9a695a add:easy dynamiCrafterLoader 2024-03-30 11:23:57 +08:00
yolain 83ca2f8ab6 add:easy ipadapterApplyEncoder and ipadapterApplyEmbeds 2024-03-29 18:00:40 +08:00
yolain 514f93c892 fix:grow_mask_by not used in easy kSamplerInpainting when additional is not none 2024-03-29 12:06:15 +08:00
yolain 9d64ea448b add:easy preMaskDetailerFix 2024-03-28 23:50:13 +08:00
yolain 1ffecfcc7d fix:faceId did not use tiled in easy ipadapterApply 2024-03-28 22:06:18 +08:00
yolain aadefaf40b add:easy ipadapterApply and easy ipadapterApplyADV 2024-03-28 21:33:47 +08:00
yolain 4c25580295 fix:inpaint patch can not use additional differential diffusion in easy kSamplerInpainting #104 2024-03-27 23:56:56 +08:00
yolain 192181007a add:differential diffusion for easy kSampelrInpainting 2024-03-25 19:19:08 +08:00
yolain 54ea3af9b8 fix:easy pipeEdit error when add lora to prompt 2024-03-25 15:25:31 +08:00
yolain c0ca3c7e95 fix:layerDiffuse xyplot bug 2024-03-25 01:34:20 +08:00
yolain 5c8af8f0b4 fix:easy pipeEdit should use the pipe steps rather than find near steps #97 2024-03-22 02:31:43 +08:00
104 changed files with 18553 additions and 1386 deletions
+21
View File
@@ -0,0 +1,21 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
paths:
- "pyproject.toml"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+5 -1
View File
@@ -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
View File
@@ -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
+114 -43
View File
@@ -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">
[![ComfyUI-Yolain-Workflows](https://github.com/yolain/ComfyUI-Easy-Use/assets/73304135/9a3f54bc-a677-4bf1-a196-8845dd57c942)](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
View File
@@ -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
View File
@@ -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
+34
View File
@@ -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)
+198 -3
View File
@@ -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
View File
@@ -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"
}
}
+332
View File
@@ -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,)
View File
+102
View File
@@ -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]}
+94
View File
@@ -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)
+95
View File
@@ -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)
)
+76
View File
@@ -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)
+219
View File
@@ -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
+762
View File
@@ -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
+809
View File
@@ -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
+143
View File
@@ -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
+82
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
View File
+156
View File
@@ -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
+23
View File
@@ -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
+167
View File
@@ -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
+185
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+10 -7
View File
@@ -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:
+5 -6
View File
@@ -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
+113
View File
@@ -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])
+13 -14
View File
@@ -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
View File
@@ -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}")
+52
View File
@@ -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({})
+115
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+27
View File
@@ -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)
+2
View File
@@ -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):
+194
View File
@@ -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()
View File
+142 -14
View File
@@ -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
View File
+58
View File
@@ -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
View File
@@ -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),)
+201
View File
@@ -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})
+148
View File
@@ -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
+238
View File
@@ -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
View File
@@ -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()
+1 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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"
}
-38
View File
@@ -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
+15
View File
@@ -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
View File
@@ -1,2 +1,5 @@
diffusers>=0.25.0
aiohttp
clip_interrogator>=0.6.0
sentencepiece==0.2.0
lark-parser
onnxruntime
+93
View File
@@ -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);
}
+35
View File
@@ -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;
}
+20
View File
@@ -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;
}
+19
View File
@@ -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
View File
@@ -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";
+265
View File
@@ -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;
}
+9 -2
View File
@@ -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
View File
@@ -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;
}
+110
View File
@@ -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;
}
+213
View File
@@ -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
View File
@@ -1,4 +1,4 @@
import { app } from "/scripts/app.js";
import { app } from "../../../scripts/app.js";
app.registerExtension({
+87 -5
View File
@@ -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
}
+5
View File
@@ -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>`
+683
View File
@@ -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;
}
}
+127
View File
@@ -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 = "&nbsp;|&nbsp;";
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
View File
@@ -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
View File
@@ -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")) {
+283
View File
@@ -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
View File
@@ -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;
},
})
+297 -68
View File
@@ -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)
+94 -3
View File
@@ -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)
}
}
});
+11 -5
View File
@@ -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) {
+3 -3
View File
@@ -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"
+2 -2
View File
@@ -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"]
+136
View File
@@ -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
View File
@@ -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 => {
+523
View File
@@ -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;
};
}
})
+2 -3
View File
@@ -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 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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++){
+237
View File
@@ -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