commit d1cd28168c3b9eff6525fac91a3d87c062e3756e Author: 刘雪峰 Date: Sat Jan 11 18:00:25 2025 +0800 init diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml new file mode 100644 index 0000000..4c358ee --- /dev/null +++ b/.github/workflows/publish.yml @@ -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 }} diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..82f9275 --- /dev/null +++ b/.gitignore @@ -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/ diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..4599e07 --- /dev/null +++ b/LICENSE @@ -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. \ No newline at end of file diff --git a/README.md b/README.md new file mode 100644 index 0000000..7c9208c --- /dev/null +++ b/README.md @@ -0,0 +1,30 @@ +[中文文档](README_CN.md) + +Add some hooks method support. Such as `TeaCache`, `PuLID-Flux`. + +## Preview (Image with WorkFlow) +![save api extended](example/workflow_base.png) + +Working with `PuLID` (need my other custom nodes [ComfyUI_PuLID_Flux_ll](https://github.com/lldacing/ComfyUI_PuLID_Flux_ll)) +![save api extended](example/PuLID_with_teacache.png) + + +## 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) diff --git a/README_CN.md b/README_CN.md new file mode 100644 index 0000000..2983403 --- /dev/null +++ b/README_CN.md @@ -0,0 +1,31 @@ +[English](README.md) + +添加一些钩子方法支持。例如支持`TeaCache`和`PulID-Flux`。 + +## 预览 (图片含工作流) +![save api extended](example/workflow_base.png) + +Working with `PuLID` (need my other custom nodes [ComfyUI_PuLID_Flux_ll](https://github.com/lldacing/ComfyUI_PuLID_Flux_ll)) +![save api extended](example/PuLID_with_teacache.png) + + +## 安装 + +- 手动安装 +```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) + diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..7704a9a --- /dev/null +++ b/__init__.py @@ -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'] diff --git a/example/PuLID_with_teacache.png b/example/PuLID_with_teacache.png new file mode 100644 index 0000000..ef0d2d9 Binary files /dev/null and b/example/PuLID_with_teacache.png differ diff --git a/example/workflow_base.png b/example/workflow_base.png new file mode 100644 index 0000000..4b441c4 Binary files /dev/null and b/example/workflow_base.png differ diff --git a/nodes/FluxPatchNode.py b/nodes/FluxPatchNode.py new file mode 100644 index 0000000..ae9c87e --- /dev/null +++ b/nodes/FluxPatchNode.py @@ -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", +} diff --git a/nodes/TeaCacheNode.py b/nodes/TeaCacheNode.py new file mode 100644 index 0000000..5a6a04c --- /dev/null +++ b/nodes/TeaCacheNode.py @@ -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", +} diff --git a/nodes/__init__.py b/nodes/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/nodes/patch_util.py b/nodes/patch_util.py new file mode 100644 index 0000000..78c56ea --- /dev/null +++ b/nodes/patch_util.py @@ -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] \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..3c42e06 --- /dev/null +++ b/pyproject.toml @@ -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 = "" diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..296d654 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +numpy \ No newline at end of file