init
This commit is contained in:
@@ -0,0 +1,23 @@
|
||||
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
|
||||
# if this is a forked repository. Skipping the workflow.
|
||||
if: github.event.repository.fork == false
|
||||
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 }}
|
||||
+162
@@ -0,0 +1,162 @@
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
|
||||
# C extensions
|
||||
*.so
|
||||
|
||||
# Distribution / packaging
|
||||
.Python
|
||||
build/
|
||||
develop-eggs/
|
||||
dist/
|
||||
downloads/
|
||||
eggs/
|
||||
.eggs/
|
||||
lib/
|
||||
lib64/
|
||||
parts/
|
||||
sdist/
|
||||
var/
|
||||
wheels/
|
||||
share/python-wheels/
|
||||
*.egg-info/
|
||||
.installed.cfg
|
||||
*.egg
|
||||
MANIFEST
|
||||
|
||||
# PyInstaller
|
||||
# Usually these files are written by a python script from a template
|
||||
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
||||
*.manifest
|
||||
*.spec
|
||||
|
||||
# Installer logs
|
||||
pip-log.txt
|
||||
pip-delete-this-directory.txt
|
||||
|
||||
# Unit test / coverage reports
|
||||
htmlcov/
|
||||
.tox/
|
||||
.nox/
|
||||
.coverage
|
||||
.coverage.*
|
||||
.cache
|
||||
nosetests.xml
|
||||
coverage.xml
|
||||
*.cover
|
||||
*.py,cover
|
||||
.hypothesis/
|
||||
.pytest_cache/
|
||||
cover/
|
||||
|
||||
# Translations
|
||||
*.mo
|
||||
*.pot
|
||||
|
||||
# Django stuff:
|
||||
*.log
|
||||
local_settings.py
|
||||
db.sqlite3
|
||||
db.sqlite3-journal
|
||||
|
||||
# Flask stuff:
|
||||
instance/
|
||||
.webassets-cache
|
||||
|
||||
# Scrapy stuff:
|
||||
.scrapy
|
||||
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
|
||||
# PyBuilder
|
||||
.pybuilder/
|
||||
target/
|
||||
|
||||
# Jupyter Notebook
|
||||
.ipynb_checkpoints
|
||||
|
||||
# IPython
|
||||
profile_default/
|
||||
ipython_config.py
|
||||
|
||||
# pyenv
|
||||
# For a library or package, you might want to ignore these files since the code is
|
||||
# intended to run in multiple environments; otherwise, check them in:
|
||||
# .python-version
|
||||
|
||||
# pipenv
|
||||
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
||||
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
||||
# having no cross-platform support, pipenv may install dependencies that don't work, or not
|
||||
# install all needed dependencies.
|
||||
#Pipfile.lock
|
||||
|
||||
# poetry
|
||||
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
|
||||
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
||||
# commonly ignored for libraries.
|
||||
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
|
||||
#poetry.lock
|
||||
|
||||
# pdm
|
||||
# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
|
||||
#pdm.lock
|
||||
# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it
|
||||
# in version control.
|
||||
# https://pdm.fming.dev/latest/usage/project/#working-with-version-control
|
||||
.pdm.toml
|
||||
.pdm-python
|
||||
.pdm-build/
|
||||
|
||||
# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
|
||||
__pypackages__/
|
||||
|
||||
# Celery stuff
|
||||
celerybeat-schedule
|
||||
celerybeat.pid
|
||||
|
||||
# SageMath parsed files
|
||||
*.sage.py
|
||||
|
||||
# Environments
|
||||
.env
|
||||
.venv
|
||||
env/
|
||||
venv/
|
||||
ENV/
|
||||
env.bak/
|
||||
venv.bak/
|
||||
|
||||
# Spyder project settings
|
||||
.spyderproject
|
||||
.spyproject
|
||||
|
||||
# Rope project settings
|
||||
.ropeproject
|
||||
|
||||
# mkdocs documentation
|
||||
/site
|
||||
|
||||
# mypy
|
||||
.mypy_cache/
|
||||
.dmypy.json
|
||||
dmypy.json
|
||||
|
||||
# Pyre type checker
|
||||
.pyre/
|
||||
|
||||
# pytype static type analyzer
|
||||
.pytype/
|
||||
|
||||
# Cython debug symbols
|
||||
cython_debug/
|
||||
|
||||
# PyCharm
|
||||
# JetBrains specific template is maintained in a separate JetBrains.gitignore that can
|
||||
# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
|
||||
# and can be added to the global gitignore or merged into this file. For a more nuclear
|
||||
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
|
||||
#.idea/
|
||||
@@ -0,0 +1,25 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2024 lldacing
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
|
||||
---
|
||||
|
||||
The code and models of BiRefNet are released under the MIT License.
|
||||
@@ -0,0 +1,30 @@
|
||||
[中文文档](README_CN.md)
|
||||
|
||||
Add some hooks method support. Such as `TeaCache`, `PuLID-Flux`.
|
||||
|
||||
## Preview (Image with WorkFlow)
|
||||

|
||||
|
||||
Working with `PuLID` (need my other custom nodes [ComfyUI_PuLID_Flux_ll](https://github.com/lldacing/ComfyUI_PuLID_Flux_ll))
|
||||

|
||||
|
||||
|
||||
## Install
|
||||
|
||||
- Manual
|
||||
```shell
|
||||
cd custom_nodes
|
||||
git clone https://github.com/lldacing/ComfyUI_Patches_ll.git
|
||||
cd ComfyUI_Patches_ll
|
||||
# restart ComfyUI
|
||||
```
|
||||
|
||||
## Nodes
|
||||
- FluxForwardOverrider
|
||||
- Add some hooks method support to the `Flux` model
|
||||
- ApplyTeaCachePatch
|
||||
- Use the `hooks` provided in `FluxForwardOverrider` to support `TeaCache` acceleration (currently only supports Flux, video related will be added in future)
|
||||
|
||||
## Thanks
|
||||
|
||||
[TeaCache](https://github.com/ali-vilab/TeaCache)
|
||||
@@ -0,0 +1,31 @@
|
||||
[English](README.md)
|
||||
|
||||
添加一些钩子方法支持。例如支持`TeaCache`和`PulID-Flux`。
|
||||
|
||||
## 预览 (图片含工作流)
|
||||

|
||||
|
||||
Working with `PuLID` (need my other custom nodes [ComfyUI_PuLID_Flux_ll](https://github.com/lldacing/ComfyUI_PuLID_Flux_ll))
|
||||

|
||||
|
||||
|
||||
## 安装
|
||||
|
||||
- 手动安装
|
||||
```shell
|
||||
cd custom_nodes
|
||||
git clone https://github.com/lldacing/ComfyUI_Patches_ll.git
|
||||
cd ComfyUI_Patches_ll
|
||||
# restart ComfyUI
|
||||
```
|
||||
|
||||
## 节点
|
||||
- FluxForwardOverrider
|
||||
- 为`Flux`模型增加一些`hook`方法支持
|
||||
- ApplyTeaCachePatch
|
||||
- 使用`FluxForwardOverrider`中的预留的hook,支持`TeaCache`加速(目前支持`Flux`,后面会加视频相关)
|
||||
|
||||
## 感谢
|
||||
|
||||
[TeaCache](https://github.com/ali-vilab/TeaCache)
|
||||
|
||||
+26
@@ -0,0 +1,26 @@
|
||||
import glob
|
||||
import importlib.util
|
||||
import os
|
||||
|
||||
extension_folder = os.path.dirname(os.path.realpath(__file__))
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
pyPath = os.path.join(extension_folder, 'nodes')
|
||||
|
||||
def loadCustomNodes():
|
||||
files = glob.glob(os.path.join(pyPath, "*Node.py"), recursive=True)
|
||||
for file in files:
|
||||
file_relative_path = file[len(extension_folder):]
|
||||
model_name = file_relative_path.replace(os.sep, '.')
|
||||
model_name = os.path.splitext(model_name)[0]
|
||||
module = importlib.import_module(model_name, __name__)
|
||||
if hasattr(module, "NODE_CLASS_MAPPINGS") and getattr(module, "NODE_CLASS_MAPPINGS") is not None:
|
||||
NODE_CLASS_MAPPINGS.update(module.NODE_CLASS_MAPPINGS)
|
||||
if hasattr(module, "NODE_DISPLAY_NAME_MAPPINGS") and getattr(module, "NODE_DISPLAY_NAME_MAPPINGS") is not None:
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(module.NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
loadCustomNodes()
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1.8 MiB |
Binary file not shown.
|
After Width: | Height: | Size: 1.7 MiB |
@@ -0,0 +1,304 @@
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
import comfy
|
||||
from .patch_util import PatchKeys
|
||||
from comfy.ldm.flux.layers import timestep_embedding
|
||||
|
||||
def flux_forward_orig(
|
||||
self,
|
||||
img: Tensor,
|
||||
img_ids: Tensor,
|
||||
txt: Tensor,
|
||||
txt_ids: Tensor,
|
||||
timesteps: Tensor,
|
||||
y: Tensor,
|
||||
guidance: Tensor = None,
|
||||
control = None,
|
||||
transformer_options={},
|
||||
attn_mask: Tensor = None,
|
||||
) -> Tensor:
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches_point = transformer_options.get(PatchKeys.options_key, {})
|
||||
|
||||
if img.ndim != 3 or txt.ndim != 3:
|
||||
raise ValueError("Input img and txt tensors must have 3 dimensions.")
|
||||
|
||||
transformer_options[PatchKeys.running_net_model] = self
|
||||
|
||||
patches_enter = patches_point.get(PatchKeys.dit_enter, [])
|
||||
if patches_enter is not None and len(patches_enter) > 0:
|
||||
for patch_enter in patches_enter:
|
||||
img, img_ids, txt, txt_ids, timesteps, y, guidance, control, attn_mask = patch_enter(img,
|
||||
img_ids,
|
||||
txt,
|
||||
txt_ids,
|
||||
timesteps,
|
||||
y,
|
||||
guidance,
|
||||
control,
|
||||
attn_mask,
|
||||
transformer_options
|
||||
)
|
||||
|
||||
# running on sequences img
|
||||
img = self.img_in(img)
|
||||
vec = self.time_in(timestep_embedding(timesteps, 256).to(img.dtype))
|
||||
if self.params.guidance_embed:
|
||||
if guidance is None:
|
||||
raise ValueError("Didn't get guidance strength for guidance distilled model.")
|
||||
vec = vec + self.guidance_in(timestep_embedding(guidance, 256).to(img.dtype))
|
||||
|
||||
vec = vec + self.vector_in(y)
|
||||
txt = self.txt_in(txt)
|
||||
|
||||
ids = torch.cat((txt_ids, img_ids), dim=1)
|
||||
pe = self.pe_embedder(ids)
|
||||
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
|
||||
patch_blocks_before = patches_point.get(PatchKeys.dit_blocks_before, [])
|
||||
if patch_blocks_before is not None and len(patch_blocks_before) > 0:
|
||||
for blocks_before in patch_blocks_before:
|
||||
img, txt, vec, ids, pe = blocks_before(img, txt, vec, ids, pe, transformer_options)
|
||||
|
||||
def double_blocks_wrap(img, txt, vec, pe, control=None, attn_mask=None, transformer_options={}):
|
||||
running_net_model = transformer_options["running_net_model"]
|
||||
for i, block in enumerate(running_net_model.double_blocks):
|
||||
# 0 -> 18
|
||||
if ("double_block", i) in blocks_replace:
|
||||
def block_wrap(args):
|
||||
out = {}
|
||||
out["img"], out["txt"] = block(img=args["img"],
|
||||
txt=args["txt"],
|
||||
vec=args["vec"],
|
||||
pe=args["pe"],
|
||||
attn_mask=args.get("attn_mask"))
|
||||
return out
|
||||
|
||||
out = blocks_replace[("double_block", i)]({"img": img,
|
||||
"txt": txt,
|
||||
"vec": vec,
|
||||
"pe": pe,
|
||||
"attn_mask": attn_mask
|
||||
},
|
||||
{
|
||||
"original_block": block_wrap,
|
||||
"transformer_options": transformer_options
|
||||
})
|
||||
txt = out["txt"]
|
||||
img = out["img"]
|
||||
else:
|
||||
img, txt = block(img=img, txt=txt, vec=vec, pe=pe, attn_mask=attn_mask)
|
||||
|
||||
if control is not None: # Controlnet
|
||||
control_i = control.get("input")
|
||||
if i < len(control_i):
|
||||
add = control_i[i]
|
||||
if add is not None:
|
||||
img += add
|
||||
|
||||
return img, txt
|
||||
|
||||
patch_double_blocks_replace = patches_point.get(PatchKeys.dit_double_blocks_replace)
|
||||
|
||||
if patch_double_blocks_replace is not None:
|
||||
img, txt = patch_double_blocks_replace({"img": img,
|
||||
"txt": txt,
|
||||
"vec": vec,
|
||||
"pe": pe,
|
||||
"control": control,
|
||||
"attn_mask": attn_mask,
|
||||
},
|
||||
{
|
||||
"original_blocks": double_blocks_wrap,
|
||||
"transformer_options": transformer_options
|
||||
})
|
||||
else:
|
||||
img, txt = double_blocks_wrap(img=img,
|
||||
txt=txt,
|
||||
vec=vec,
|
||||
pe=pe,
|
||||
control=control,
|
||||
attn_mask=attn_mask,
|
||||
transformer_options=transformer_options
|
||||
)
|
||||
|
||||
patches_double_blocks_after = patches_point.get(PatchKeys.dit_double_blocks_after, [])
|
||||
if patches_double_blocks_after is not None and len(patches_double_blocks_after) > 0:
|
||||
for patch_double_blocks_after in patches_double_blocks_after:
|
||||
img, txt = patch_double_blocks_after(img, txt, transformer_options)
|
||||
|
||||
patch_blocks_transition = patches_point.get(PatchKeys.dit_blocks_transition_replace)
|
||||
|
||||
def blocks_transition_wrap(**kwargs):
|
||||
txt = kwargs["txt"]
|
||||
img = kwargs["img"]
|
||||
return torch.cat((txt, img), 1)
|
||||
|
||||
if patch_blocks_transition is not None:
|
||||
img = patch_blocks_transition({"img": img, "txt": txt, "vec": vec, "pe": pe},
|
||||
{
|
||||
"original_func": blocks_transition_wrap,
|
||||
"transformer_options": transformer_options
|
||||
})
|
||||
else:
|
||||
img = blocks_transition_wrap(img=img, txt=txt)
|
||||
|
||||
patches_single_blocks_before = patches_point.get(PatchKeys.dit_single_blocks_before, [])
|
||||
if patches_single_blocks_before is not None and len(patches_single_blocks_before) > 0:
|
||||
for patch_single_blocks_before in patches_single_blocks_before:
|
||||
img, txt = patch_single_blocks_before(img, txt, transformer_options)
|
||||
|
||||
def single_blocks_wrap(img, txt, vec, pe, control=None, attn_mask=None, transformer_options={}):
|
||||
running_net_model = transformer_options[PatchKeys.running_net_model]
|
||||
for i, block in enumerate(running_net_model.single_blocks):
|
||||
# 0 -> 37
|
||||
if ("single_block", i) in blocks_replace:
|
||||
def block_wrap(args):
|
||||
out = {}
|
||||
out["img"] = block(args["img"],
|
||||
vec=args["vec"],
|
||||
pe=args["pe"],
|
||||
attn_mask=args.get("attn_mask"))
|
||||
return out
|
||||
|
||||
out = blocks_replace[("single_block", i)]({"img": img,
|
||||
"vec": vec,
|
||||
"pe": pe,
|
||||
"attn_mask": attn_mask},
|
||||
{
|
||||
"original_block": block_wrap,
|
||||
"transformer_options": transformer_options
|
||||
})
|
||||
img = out["img"]
|
||||
else:
|
||||
img = block(img, vec=vec, pe=pe, attn_mask=attn_mask)
|
||||
|
||||
if control is not None: # Controlnet
|
||||
control_o = control.get("output")
|
||||
if i < len(control_o):
|
||||
add = control_o[i]
|
||||
if add is not None:
|
||||
img[:, txt.shape[1]:, ...] += add
|
||||
|
||||
return img
|
||||
|
||||
patch_single_blocks_replace = patches_point.get(PatchKeys.dit_single_blocks_replace)
|
||||
|
||||
if patch_single_blocks_replace is not None:
|
||||
img, txt = patch_single_blocks_replace({"img": img,
|
||||
"txt": txt,
|
||||
"vec": vec,
|
||||
"pe": pe,
|
||||
"control": control,
|
||||
"attn_mask": attn_mask
|
||||
},
|
||||
{
|
||||
"original_blocks": single_blocks_wrap,
|
||||
"transformer_options": transformer_options
|
||||
})
|
||||
else:
|
||||
img = single_blocks_wrap(img=img,
|
||||
txt=txt,
|
||||
vec=vec,
|
||||
pe=pe,
|
||||
control=control,
|
||||
attn_mask=attn_mask,
|
||||
transformer_options=transformer_options
|
||||
)
|
||||
|
||||
patch_blocks_exit = patches_point.get(PatchKeys.dit_blocks_after, [])
|
||||
if patch_blocks_exit is not None and len(patch_blocks_exit) > 0:
|
||||
for blocks_after in patch_blocks_exit:
|
||||
img, txt = blocks_after(img, txt, transformer_options)
|
||||
|
||||
def final_transition_wrap(**kwargs):
|
||||
img = kwargs["img"]
|
||||
txt = kwargs["txt"]
|
||||
return img[:, txt.shape[1]:, ...]
|
||||
|
||||
patch_blocks_after_transition_replace = patches_point.get(PatchKeys.dit_blocks_after_transition_replace)
|
||||
if patch_blocks_after_transition_replace is not None:
|
||||
img = patch_blocks_after_transition_replace({"img": img, "txt": txt, "vec": vec, "pe": pe},
|
||||
{
|
||||
"original_func": final_transition_wrap,
|
||||
"transformer_options": transformer_options
|
||||
})
|
||||
else:
|
||||
img = final_transition_wrap(img=img, txt=txt)
|
||||
|
||||
patches_final_layer_before = patches_point.get(PatchKeys.dit_final_layer_before, [])
|
||||
if patches_final_layer_before is not None and len(patches_final_layer_before) > 0:
|
||||
for patch_final_layer_before in patches_final_layer_before:
|
||||
img = patch_final_layer_before(img, txt, transformer_options)
|
||||
|
||||
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
|
||||
|
||||
patches_exit = patches_point.get(PatchKeys.dit_exit, [])
|
||||
if patches_exit is not None and len(patches_exit) > 0:
|
||||
for patch_exit in patches_exit:
|
||||
img = patch_exit(img, transformer_options)
|
||||
|
||||
del transformer_options[PatchKeys.running_net_model]
|
||||
|
||||
return img
|
||||
|
||||
|
||||
def outer_sample_function_wrapper(wrapper_executor, noise, latent_image, sampler, sigmas, denoise_mask=None,
|
||||
callback=None, disable_pbar=False, seed=None):
|
||||
# set hook
|
||||
set_hook()
|
||||
|
||||
try:
|
||||
out = wrapper_executor(noise, latent_image, sampler, sigmas, denoise_mask=denoise_mask, callback=callback,
|
||||
disable_pbar=disable_pbar, seed=seed)
|
||||
finally:
|
||||
# cleanup hook
|
||||
clean_hook()
|
||||
return out
|
||||
|
||||
def set_hook():
|
||||
comfy.ldm.flux.model.Flux.class_old_forward_orig = comfy.ldm.flux.model.Flux.forward_orig
|
||||
comfy.ldm.flux.model.Flux.forward_orig = flux_forward_orig
|
||||
|
||||
def clean_hook():
|
||||
if hasattr(comfy.ldm.flux.model.Flux, 'class_old_forward_orig'):
|
||||
comfy.ldm.flux.model.Flux.forward_orig = comfy.ldm.flux.model.Flux.class_old_forward_orig
|
||||
del comfy.ldm.flux.model.Flux.class_old_forward_orig
|
||||
|
||||
class FluxForwardOverrider:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
RETURN_NAMES = ("model",)
|
||||
FUNCTION = "apply_patch"
|
||||
CATEGORY = "patches/flux"
|
||||
|
||||
def apply_patch(self, model):
|
||||
|
||||
model = model.clone()
|
||||
patch_key = "flux_forward_override_wrapper"
|
||||
if len(model.get_wrappers(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, patch_key)) == 0:
|
||||
# Just add it once when connecting in series
|
||||
model.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE,
|
||||
patch_key,
|
||||
outer_sample_function_wrapper
|
||||
)
|
||||
return (model, )
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"FluxForwardOverrider": FluxForwardOverrider,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FluxForwardOverrider": "FluxForwardOverrider",
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
import numpy as np
|
||||
|
||||
import comfy
|
||||
from .patch_util import PatchKeys, add_model_patch_option, set_model_patch, set_model_patch_replace
|
||||
|
||||
tea_cache_key_attrs = "tea_cache_attr"
|
||||
|
||||
def tea_cache_enter(img, img_ids, txt, txt_ids, timesteps, y, guidance, control, attn_mask, transformer_options):
|
||||
diffusion_model = transformer_options.get(PatchKeys.running_net_model)
|
||||
if hasattr(diffusion_model, "flux_tea_cache"):
|
||||
tea_cache = getattr(diffusion_model, "flux_tea_cache", {})
|
||||
transformer_options[tea_cache_key_attrs] = tea_cache
|
||||
return img, img_ids, txt, txt_ids, timesteps, y, guidance, control, attn_mask
|
||||
|
||||
def tea_cache_patch_blocks_before(img, txt, vec, ids, pe, transformer_options):
|
||||
real_model = transformer_options[PatchKeys.running_net_model]
|
||||
attrs = transformer_options.get(tea_cache_key_attrs, {})
|
||||
|
||||
# tea cache src code
|
||||
# if self.emb is not None:
|
||||
# emb = self.emb(timestep, class_labels, hidden_dtype=hidden_dtype)
|
||||
# emb = self.linear(self.silu(emb))
|
||||
# shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = emb.chunk(6, dim=1)
|
||||
# x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None]
|
||||
# x, gate_msa, shift_mlp, scale_mlp, gate_mlp
|
||||
inp = img.clone()
|
||||
vec_ = vec.clone()
|
||||
double_block_0 = real_model.double_blocks[0]
|
||||
img_mod1, img_mod2 = double_block_0.img_mod(vec_)
|
||||
modulated_inp = double_block_0.img_norm1(inp)
|
||||
modulated_inp = (1 + img_mod1.scale) * modulated_inp + img_mod1.shift
|
||||
if attrs['cnt'] == 0 or attrs['cnt'] == attrs['total_steps'] - 1:
|
||||
should_calc = True
|
||||
attrs['accumulated_rel_l1_distance'] = 0
|
||||
else:
|
||||
coefficients = [4.98651651e+02, -2.83781631e+02, 5.58554382e+01, -3.82021401e+00, 2.64230861e-01]
|
||||
rescale_func = np.poly1d(coefficients)
|
||||
attrs['accumulated_rel_l1_distance'] += rescale_func(((modulated_inp - attrs['previous_modulated_input']).abs().mean() / attrs['previous_modulated_input'].abs().mean()).cpu().item())
|
||||
|
||||
if attrs['accumulated_rel_l1_distance'] < attrs['rel_l1_thresh']:
|
||||
should_calc = False
|
||||
else:
|
||||
should_calc = True
|
||||
attrs['accumulated_rel_l1_distance'] = 0
|
||||
attrs['previous_modulated_input'] = modulated_inp
|
||||
attrs['cnt'] += 1
|
||||
if attrs['cnt'] == attrs['total_steps']:
|
||||
attrs['cnt'] = 0
|
||||
|
||||
attrs['should_calc'] = should_calc
|
||||
|
||||
return img, txt, vec, ids, pe
|
||||
|
||||
def tea_cache_patch_double_blocks_replace(original_args, wrapper_options):
|
||||
img = original_args['img']
|
||||
txt = original_args['txt']
|
||||
transformer_options = wrapper_options.get('transformer_options', {})
|
||||
attrs = transformer_options.get(tea_cache_key_attrs, {})
|
||||
should_calc = attrs.get('should_calc', True)
|
||||
if not should_calc:
|
||||
img += attrs['previous_residual']
|
||||
else:
|
||||
# (b, seq_len, _)
|
||||
attrs['ori_img'] = img.clone()
|
||||
img, txt = wrapper_options.get('original_blocks')(**original_args, transformer_options=transformer_options)
|
||||
return img, txt
|
||||
|
||||
def tea_cache_patch_blocks_transition_replace(original_args, wrapper_options):
|
||||
img = original_args['img']
|
||||
transformer_options = wrapper_options.get('transformer_options', {})
|
||||
attrs = transformer_options.get(tea_cache_key_attrs, {})
|
||||
should_calc = attrs.get('should_calc', True)
|
||||
if should_calc:
|
||||
img = wrapper_options.get('original_func')(**original_args, transformer_options=transformer_options)
|
||||
return img
|
||||
|
||||
def tea_cache_patch_single_blocks_replace(original_args, wrapper_options):
|
||||
img = original_args['img']
|
||||
txt = original_args['txt']
|
||||
transformer_options = wrapper_options.get('transformer_options', {})
|
||||
attrs = transformer_options.get(tea_cache_key_attrs, {})
|
||||
should_calc = attrs.get('should_calc', True)
|
||||
if should_calc:
|
||||
img = wrapper_options.get('original_blocks')(**original_args, transformer_options=transformer_options)
|
||||
return img, txt
|
||||
|
||||
def tea_cache_patch_blocks_after_replace(original_args, wrapper_options):
|
||||
img = original_args['img']
|
||||
transformer_options = wrapper_options.get('transformer_options', {})
|
||||
attrs = transformer_options.get(tea_cache_key_attrs, {})
|
||||
should_calc = attrs.get('should_calc', True)
|
||||
if should_calc:
|
||||
img = wrapper_options.get('original_func')(**original_args)
|
||||
return img
|
||||
|
||||
def tea_cache_patch_final_transition_after(img, txt, transformer_options):
|
||||
attrs = transformer_options.get(tea_cache_key_attrs, {})
|
||||
should_calc = attrs.get('should_calc', True)
|
||||
if should_calc:
|
||||
attrs['previous_residual'] = img - attrs['ori_img']
|
||||
return img
|
||||
|
||||
def tea_cache_patch_dit_exit(img, transformer_options):
|
||||
tea_cache = transformer_options.get(tea_cache_key_attrs, {})
|
||||
setattr(transformer_options.get(PatchKeys.running_net_model), "flux_tea_cache", tea_cache)
|
||||
return img
|
||||
|
||||
def tea_cache_prepare_wrapper(wrapper_executor, noise, latent_image, sampler, sigmas, denoise_mask=None,
|
||||
callback=None, disable_pbar=False, seed=None):
|
||||
cfg_guider = wrapper_executor.class_obj
|
||||
|
||||
# Use cfd_guider.model_options, which is copied from modelPatcher.model_options and will be restored after execution without any unexpected contamination
|
||||
temp_options = add_model_patch_option(cfg_guider, tea_cache_key_attrs)
|
||||
temp_options['total_steps'] = len(sigmas) - 1
|
||||
temp_options['cnt'] = 0
|
||||
try:
|
||||
out = wrapper_executor(noise, latent_image, sampler, sigmas, denoise_mask=denoise_mask, callback=callback,
|
||||
disable_pbar=disable_pbar, seed=seed)
|
||||
finally:
|
||||
diffusion_model = cfg_guider.model_patcher.model.diffusion_model
|
||||
if hasattr(diffusion_model, "flux_tea_cache"):
|
||||
del diffusion_model.flux_tea_cache
|
||||
|
||||
return out
|
||||
|
||||
class ApplyTeaCachePatch:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"rel_l1_thresh": ("FLOAT",
|
||||
{
|
||||
"default": 0.25,
|
||||
"min": 0.0,
|
||||
"max": 5.0,
|
||||
"step": 0.01,
|
||||
"tooltip": "0 (original), 0.25 (1.5x speedup), 0.4 (1.8x speedup), 0.6 (2.0x speedup), and 0.8 (2.25x speedup)."
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
RETURN_NAMES = ("model",)
|
||||
FUNCTION = "apply_patch"
|
||||
CATEGORY = "patches/flux"
|
||||
DESCRIPTION = "TeaCache加速补丁"
|
||||
|
||||
def apply_patch(self, model, rel_l1_thresh):
|
||||
|
||||
model = model.clone()
|
||||
|
||||
set_model_patch(model, PatchKeys.options_key, tea_cache_enter, PatchKeys.dit_enter)
|
||||
set_model_patch(model, PatchKeys.options_key, tea_cache_patch_blocks_before, PatchKeys.dit_blocks_before)
|
||||
|
||||
set_model_patch_replace(model, PatchKeys.options_key, tea_cache_patch_double_blocks_replace, PatchKeys.dit_double_blocks_replace)
|
||||
set_model_patch_replace(model, PatchKeys.options_key, tea_cache_patch_blocks_transition_replace, PatchKeys.dit_blocks_transition_replace)
|
||||
set_model_patch_replace(model, PatchKeys.options_key, tea_cache_patch_single_blocks_replace, PatchKeys.dit_single_blocks_replace)
|
||||
set_model_patch_replace(model, PatchKeys.options_key, tea_cache_patch_blocks_after_replace, PatchKeys.dit_blocks_after_transition_replace)
|
||||
|
||||
set_model_patch(model, PatchKeys.options_key, tea_cache_patch_final_transition_after, PatchKeys.dit_final_layer_before)
|
||||
set_model_patch(model, PatchKeys.options_key, tea_cache_patch_dit_exit, PatchKeys.dit_exit)
|
||||
|
||||
flux_forward_patch = add_model_patch_option(model, tea_cache_key_attrs)
|
||||
flux_forward_patch['rel_l1_thresh'] = rel_l1_thresh
|
||||
|
||||
patch_key = "tea_cache_wrapper"
|
||||
if len(model.get_wrappers(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, patch_key)) == 0:
|
||||
# Just add it once when connecting in series
|
||||
model.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE,
|
||||
patch_key,
|
||||
tea_cache_prepare_wrapper
|
||||
)
|
||||
return (model, )
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"ApplyTeaCachePatch": ApplyTeaCachePatch,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ApplyTeaCachePatch": "ApplyTeaCachePatch",
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
class PatchKeys:
|
||||
################## transformer_options patches ##################
|
||||
options_key = "patches_point"
|
||||
running_net_model = "running_net_model"
|
||||
# patches_point下支持设置的补丁
|
||||
dit_enter = "patch_dit_enter"
|
||||
dit_blocks_before = "patch_dit_blocks_before"
|
||||
dit_double_blocks_replace = "patch_dit_double_blocks_replace"
|
||||
dit_double_blocks_after = "patch_dit_double_blocks_after"
|
||||
dit_blocks_transition_replace = "patch_dit_blocks_transition_replace"
|
||||
dit_single_blocks_before = "patch_dit_single_blocks_before"
|
||||
dit_single_blocks_replace = "patch_dit_single_blocks_replace"
|
||||
dit_blocks_after = "patch_dit_blocks_after"
|
||||
dit_blocks_after_transition_replace = "patch_dit_final_layer_before_replace"
|
||||
dit_final_layer_before = "patch_dit_final_layer_before"
|
||||
dit_exit = "patch_dit_exit"
|
||||
################## transformer_options patches ##################
|
||||
|
||||
|
||||
def set_model_patch(model_patcher, options_key, patch, name):
|
||||
to = model_patcher.model_options["transformer_options"]
|
||||
if options_key not in to:
|
||||
to[options_key] = {}
|
||||
to[options_key][name] = to[options_key].get(name, []) + [patch]
|
||||
|
||||
def set_model_patch_replace(model_patcher, options_key, patch, name):
|
||||
to = model_patcher.model_options["transformer_options"]
|
||||
if options_key not in to:
|
||||
to[options_key] = {}
|
||||
to[options_key][name] = patch
|
||||
|
||||
def add_model_patch_option(model, patch_key):
|
||||
if 'transformer_options' not in model.model_options:
|
||||
model.model_options['transformer_options'] = {}
|
||||
to = model.model_options['transformer_options']
|
||||
if patch_key not in to:
|
||||
to[patch_key] = {}
|
||||
return to[patch_key]
|
||||
@@ -0,0 +1,15 @@
|
||||
[project]
|
||||
name = "comfyui_patches_ll"
|
||||
description = "Some patches for Flux etc, support TeaCache, PuLID."
|
||||
version = "1.0.0"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = []
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/lldacing/ComfyUI_Patches_ll"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "lldacing"
|
||||
DisplayName = "ComfyUI_Patches_ll"
|
||||
Icon = ""
|
||||
@@ -0,0 +1 @@
|
||||
numpy
|
||||
Reference in New Issue
Block a user