diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml new file mode 100644 index 0000000..050989f --- /dev/null +++ b/.github/workflows/publish.yml @@ -0,0 +1,28 @@ +name: Publish to Comfy registry +on: + workflow_dispatch: + push: + branches: + - main + - master + paths: + - "pyproject.toml" + +permissions: + issues: write + +jobs: + publish-node: + name: Publish Custom Node to registry + runs-on: ubuntu-latest + if: ${{ github.repository_owner == 'wildminder' }} + steps: + - name: Check out code + uses: actions/checkout@v4 + with: + submodules: true + - name: Publish Custom Node + uses: Comfy-Org/publish-node-action@v1 + 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..0efcd5d --- /dev/null +++ b/.gitignore @@ -0,0 +1,163 @@ +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.idea +.Python +__pycache__ +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/#use-with-ide +.pdm.toml + +# 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/ +.rar + +# 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/__init__.py b/__init__.py new file mode 100644 index 0000000..12511c5 --- /dev/null +++ b/__init__.py @@ -0,0 +1,92 @@ +import torch +from comfy_api.latest import ComfyExtension, io +from .src.patch import apply_dype_to_flux + +class DyPE_FLUX(io.ComfyNode): + """ + Applies DyPE (Dynamic Position Extrapolation) to a FLUX model. + This allows generating images at resolutions far beyond the model's training scale + by dynamically adjusting positional encodings and the noise schedule. + """ + + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="DyPE_FLUX", + display_name="DyPE for FLUX", + category="model_patches/unet", + description="Applies DyPE (Dynamic Position Extrapolation) to a FLUX model for ultra-high-resolution generation.", + inputs=[ + io.Model.Input( + "model", + tooltip="The FLUX model to patch with DyPE.", + ), + io.Int.Input( + "width", + default=1024, min=16, max=8192, step=8, + tooltip="Target image width. Must match the width of your empty latent." + ), + io.Int.Input( + "height", + default=1024, min=16, max=8192, step=8, + tooltip="Target image height. Must match the height of your empty latent." + ), + io.Combo.Input( + "method", + options=["yarn", "ntk", "base"], + default="yarn", + tooltip="Position encoding extrapolation method (YARN recommended).", + ), + io.Boolean.Input( + "enable_dype", + default=True, + label_on="Enabled", + label_off="Disabled", + tooltip="Enable or disable Dynamic Position Extrapolation for RoPE.", + ), + io.Float.Input( + "dype_exponent", + default=2.0, min=0.0, max=4.0, step=0.1, + optional=True, + tooltip="Controls DyPE strength over time (λt). 2.0=Exponential (best for 4K+), 1.0=Linear, 0.5=Sub-linear (better for ~2K)." + ), + io.Float.Input( + "base_shift", + default=0.5, min=0.0, max=10.0, step=0.01, + optional=True, + tooltip="Advanced: Base shift for the noise schedule (mu). Default is 0.5." + ), + io.Float.Input( + "max_shift", + default=1.15, min=0.0, max=10.0, step=0.01, + optional=True, + tooltip="Advanced: Max shift for the noise schedule (mu) at high resolutions. Default is 1.15." + ), + ], + outputs=[ + io.Model.Output( + display_name="Patched Model", + tooltip="The FLUX model patched with DyPE.", + ), + ], + ) + + @classmethod + def execute(cls, model, width: int, height: int, method: str, enable_dype: bool, dype_exponent: float = 2.0, base_shift: float = 0.5, max_shift: float = 1.15) -> io.NodeOutput: + """ + Clones the model and applies the DyPE patch for both the noise schedule and positional embeddings. + """ + if not hasattr(model.model, "diffusion_model") or not hasattr(model.model.diffusion_model, "pe_embedder"): + raise ValueError("This node is only compatible with FLUX models.") + + patched_model = apply_dype_to_flux(model, width, height, method, enable_dype, dype_exponent, base_shift, max_shift) + return io.NodeOutput(patched_model) + +class DyPEExtension(ComfyExtension): + """Registers the DyPE node.""" + + async def get_node_list(self) -> list[type[io.ComfyNode]]: + return [DyPE_FLUX] + +async def comfy_entrypoint() -> DyPEExtension: + return DyPEExtension() \ No newline at end of file diff --git a/example_workflows/DyPE-Flux-workflow.json b/example_workflows/DyPE-Flux-workflow.json new file mode 100644 index 0000000..0474774 --- /dev/null +++ b/example_workflows/DyPE-Flux-workflow.json @@ -0,0 +1,844 @@ +{ + "id": "908d0bfb-e192-4627-9b57-147496e6e2dd", + "revision": 0, + "last_node_id": 70, + "last_link_id": 119, + "nodes": [ + { + "id": 40, + "type": "DualCLIPLoader", + "pos": [ + -320, + 290 + ], + "size": [ + 270, + 130 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "CLIP", + "type": "CLIP", + "links": [ + 64 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.40", + "Node name for S&R": "DualCLIPLoader", + "models": [ + { + "name": "clip_l.safetensors", + "url": "https://huggingface.co/comfyanonymous/flux_text_encoders/resolve/main/clip_l.safetensors", + "directory": "text_encoders" + }, + { + "name": "t5xxl_fp16.safetensors", + "url": "https://huggingface.co/comfyanonymous/flux_text_encoders/resolve/main/t5xxl_fp16.safetensors", + "directory": "text_encoders" + } + ] + }, + "widgets_values": [ + "clip_l.safetensors", + "t5xxl_fp16.safetensors", + "flux", + "default" + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 43, + "type": "MarkdownNote", + "pos": [ + -870, + 110 + ], + "size": [ + 520, + 390 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "title": "Model links", + "properties": {}, + "widgets_values": [ + "## Model links\n\n**Diffusion Model**\n\n- [flux1-krea-dev_fp8_scaled.safetensors](https://huggingface.co/Comfy-Org/FLUX.1-Krea-dev_ComfyUI/resolve/main/split_files/diffusion_models/flux1-krea-dev_fp8_scaled.safetensors)\n\nIf you need the original weights, head to [black-forest-labs/FLUX.1-Krea-dev](https://huggingface.co/black-forest-labs/FLUX.1-Krea-dev/), accept the agreement in the repo, then click the link below to download the models:\n\n- [flux1-krea-dev.safetensors](https://huggingface.co/black-forest-labs/FLUX.1-Krea-dev/resolve/main/flux1-krea-dev.safetensors)\n\n**Text Encoder**\n\n- [clip_l.safetensors](https://huggingface.co/comfyanonymous/flux_text_encoders/blob/main/clip_l.safetensors)\n\n- [t5xxl_fp16.safetensors](https://huggingface.co/comfyanonymous/flux_text_encoders/resolve/main/t5xxl_fp16.safetensors) or [t5xxl_fp8_e4m3fn_scaled.safetensors](https://huggingface.co/comfyanonymous/flux_text_encoders/resolve/main/t5xxl_fp8_e4m3fn_scaled.safetensors)\n\n**VAE**\n\n- [ae.safetensors](https://huggingface.co/Comfy-Org/Lumina_Image_2.0_Repackaged/resolve/main/split_files/vae/ae.safetensors)\n\n\n```\nComfyUI/\n├── models/\n│ ├── diffusion_models/\n│ │ └─── flux1-krea-dev_fp8_scaled.safetensors\n│ ├── text_encoders/\n│ │ ├── clip_l.safetensors\n│ │ └─── t5xxl_fp16.safetensors # or t5xxl_fp8_e4m3fn_scaled.safetensors\n│ └── vae/\n│ └── ae.safetensors\n```\n" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 39, + "type": "VAELoader", + "pos": [ + -320, + 470 + ], + "size": [ + 270, + 58 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "VAE", + "type": "VAE", + "links": [ + 77, + 85 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.40", + "Node name for S&R": "VAELoader", + "models": [ + { + "name": "ae.safetensors", + "url": "https://huggingface.co/Comfy-Org/Lumina_Image_2.0_Repackaged/resolve/main/split_files/vae/ae.safetensors", + "directory": "vae" + } + ] + }, + "widgets_values": [ + "FLUX1\\ae.safetensors" + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 8, + "type": "VAEDecode", + "pos": [ + 391.97558139331176, + 122.95473372107558 + ], + "size": [ + 319.2356335124862, + 46 + ], + "flags": { + "collapsed": false + }, + "order": 13, + "mode": 0, + "inputs": [ + { + "name": "samples", + "type": "LATENT", + "link": 52 + }, + { + "name": "vae", + "type": "VAE", + "link": 85 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "slot_index": 0, + "links": [] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.40", + "Node name for S&R": "VAEDecode" + }, + "widgets_values": [] + }, + { + "id": 55, + "type": "VAEDecodeTiled", + "pos": [ + 391.97558139331176, + 226.612614517858 + ], + "size": [ + 322.89304359551966, + 150 + ], + "flags": {}, + "order": 14, + "mode": 0, + "inputs": [ + { + "name": "samples", + "type": "LATENT", + "link": 112 + }, + { + "name": "vae", + "type": "VAE", + "link": 77 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 91 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.66", + "Node name for S&R": "VAEDecodeTiled" + }, + "widgets_values": [ + 256, + 64, + 64, + 8 + ] + }, + { + "id": 53, + "type": "ModelPatchTorchSettings", + "pos": [ + 30.020388325880088, + 958.5302639048878 + ], + "size": [ + 280, + 58 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "MODEL", + "link": 71 + } + ], + "outputs": [ + { + "name": "MODEL", + "type": "MODEL", + "links": [ + 73 + ] + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "5dcda71011870278c35d92ff77a677ed2e538f2d", + "Node name for S&R": "ModelPatchTorchSettings", + "ue_properties": { + "version": "7.0.1", + "widget_ue_connectable": {} + } + }, + "widgets_values": [ + true + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 38, + "type": "UNETLoader", + "pos": [ + -320, + 150 + ], + "size": [ + 270, + 82 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "MODEL", + "type": "MODEL", + "links": [ + 118 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.40", + "Node name for S&R": "UNETLoader", + "models": [ + { + "name": "flux1-krea-dev_fp8_scaled.safetensors", + "url": "https://huggingface.co/Comfy-Org/FLUX.1-Krea-dev_ComfyUI/resolve/main/split_files/diffusion_models/flux1-krea-dev_fp8_scaled.safetensors", + "directory": "diffusion_models" + } + ] + }, + "widgets_values": [ + "FLUX1\\flux1-krea-dev.safetensors", + "default" + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 27, + "type": "EmptySD3LatentImage", + "pos": [ + -320, + 630 + ], + "size": [ + 270, + 120 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "LATENT", + "type": "LATENT", + "slot_index": 0, + "links": [ + 51 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.40", + "Node name for S&R": "EmptySD3LatentImage" + }, + "widgets_values": [ + 2048, + 3072, + 1 + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 9, + "type": "SaveImage", + "pos": [ + 757.2214793721178, + 123.42500709283193 + ], + "size": [ + 640, + 660 + ], + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 91 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.40", + "Node name for S&R": "SaveImage" + }, + "widgets_values": [ + "flux_krea/flux_krea" + ] + }, + { + "id": 45, + "type": "CLIPTextEncode", + "pos": [ + 7.045080197257759, + 159.16529405661166 + ], + "size": [ + 330, + 210 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "clip", + "type": "CLIP", + "link": 64 + } + ], + "outputs": [ + { + "name": "CONDITIONING", + "type": "CONDITIONING", + "links": [ + 66, + 111 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.47", + "Node name for S&R": "CLIPTextEncode" + }, + "widgets_values": [ + "A muscular, bald man holds a flower above his head with both arms, set against a soft circular background in black and white." + ], + "color": "#232", + "bgcolor": "#353" + }, + { + "id": 42, + "type": "ConditioningZeroOut", + "pos": [ + 151.75751860615074, + 416.7514387263594 + ], + "size": [ + 200, + 30 + ], + "flags": { + "collapsed": true + }, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "conditioning", + "type": "CONDITIONING", + "link": 66 + } + ], + "outputs": [ + { + "name": "CONDITIONING", + "type": "CONDITIONING", + "links": [ + 110 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.40", + "Node name for S&R": "ConditioningZeroOut" + }, + "widgets_values": [] + }, + { + "id": 52, + "type": "PathchSageAttentionKJ", + "pos": [ + 30.97067656870301, + 849.9776697194811 + ], + "size": [ + 280, + 58 + ], + "flags": { + "collapsed": false + }, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "MODEL", + "link": 119 + } + ], + "outputs": [ + { + "name": "MODEL", + "type": "MODEL", + "links": [ + 71 + ] + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "5dcda71011870278c35d92ff77a677ed2e538f2d", + "Node name for S&R": "PathchSageAttentionKJ", + "ue_properties": { + "version": "7.0.1", + "widget_ue_connectable": {} + } + }, + "widgets_values": [ + "auto" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 68, + "type": "DyPE_FLUX", + "pos": [ + 34.378266361112125, + 592.0845558725733 + ], + "size": [ + 273.70012497212906, + 202 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "MODEL", + "link": 118 + } + ], + "outputs": [ + { + "name": "Patched Model", + "type": "MODEL", + "links": [ + 119 + ] + } + ], + "properties": { + "Node name for S&R": "DyPE_FLUX" + }, + "widgets_values": [ + 1024, + 1024, + "yarn", + true, + 3, + 0.1, + 0.8 + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 70, + "type": "Note", + "pos": [ + 370.8346374838202, + 962.8353927340362 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "DyPE width/height - Keep the values below 1024x1024; doing so won’t affect your output.\n" + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 69, + "type": "MarkdownNote", + "pos": [ + -866.3440753786277, + 557.8510364476188 + ], + "size": [ + 519.2421649643175, + 338.01057378012047 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [], + "title": "DyPE", + "properties": { + "ue_properties": { + "widget_ue_connectable": {}, + "version": "7.1", + "input_ue_unconnectable": {} + } + }, + "widgets_values": [ + "### Node Inputs\n\n* **`model`**: The FLUX model to be patched.\n* **`width` / `height`**: The target image resolution. **This must match the resolution set in your `Empty Latent Image` node.**\n* **`method`**: The core position encoding extrapolation method. `yarn` is the recommended default, as it forms the basis of the paper's best-performing \"DY-YaRN\" variant.\n* **`enable_dype`**: Enables or disables the **dynamic, time-aware** component of DyPE.\n* **`dype_exponent`**: Controls the \"strength\" of the dynamic effect over time. This is the most important tuning parameter.\n * `2.0` (Exponential): Recommended for **4K+** resolutions. It's an aggressive schedule that transitions quickly.\n * `1.0` (Linear): A good starting point for **~2K-3K** resolutions.\n * `0.5` (Sub-linear): A gentler schedule that may work best for resolutions just above the model's native 1K.\n* **`base_shift` / `max_shift`** (Advanced): Adjust only if you are an advanced user experimenting with the noise schedule.\n\n\n\n## Join\n\n### [TokenDiffusion](https://t.me/TokenDiff) - AI for every home, creativity for every mind!\n\n### [TokenDiff Community Hub](https://t.me/TokenDiff_hub) - Questions, help, and thoughtful discussion. " + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 31, + "type": "KSampler", + "pos": [ + 395.5510343546297, + 433.7293334706151 + ], + "size": [ + 315, + 474.00000000000006 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "MODEL", + "link": 73 + }, + { + "name": "positive", + "type": "CONDITIONING", + "link": 111 + }, + { + "name": "negative", + "type": "CONDITIONING", + "link": 110 + }, + { + "name": "latent_image", + "type": "LATENT", + "link": 51 + } + ], + "outputs": [ + { + "name": "LATENT", + "type": "LATENT", + "slot_index": 0, + "links": [ + 52, + 112 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.40", + "Node name for S&R": "KSampler" + }, + "widgets_values": [ + 42, + "fixed", + 30, + 1, + "euler", + "beta", + 1 + ] + } + ], + "links": [ + [ + 51, + 27, + 0, + 31, + 3, + "LATENT" + ], + [ + 52, + 31, + 0, + 8, + 0, + "LATENT" + ], + [ + 64, + 40, + 0, + 45, + 0, + "CLIP" + ], + [ + 66, + 45, + 0, + 42, + 0, + "CONDITIONING" + ], + [ + 71, + 52, + 0, + 53, + 0, + "MODEL" + ], + [ + 73, + 53, + 0, + 31, + 0, + "MODEL" + ], + [ + 77, + 39, + 0, + 55, + 1, + "VAE" + ], + [ + 85, + 39, + 0, + 8, + 1, + "VAE" + ], + [ + 91, + 55, + 0, + 9, + 0, + "IMAGE" + ], + [ + 110, + 42, + 0, + 31, + 2, + "CONDITIONING" + ], + [ + 111, + 45, + 0, + 31, + 1, + "CONDITIONING" + ], + [ + 112, + 31, + 0, + 55, + 0, + "LATENT" + ], + [ + 118, + 38, + 0, + 68, + 0, + "MODEL" + ], + [ + 119, + 68, + 0, + 52, + 0, + "MODEL" + ] + ], + "groups": [ + { + "id": 1, + "title": "Step 1 - Load Models Here", + "bounding": [ + -330, + 80, + 300, + 460 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Step 2 - Image Size", + "bounding": [ + -330, + 560, + 300, + 200 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + }, + { + "id": 3, + "title": "Step 3 - Prompt", + "bounding": [ + -10, + 80, + 361.9005764856456, + 399.16989485829595 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + }, + { + "id": 5, + "title": "Model patch", + "bounding": [ + -7.575982245696814, + 506.15777888556613, + 357.2928781597926, + 546.8190068170823 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.7627768444385601, + "offset": [ + 1067.6762292181238, + -14.895608507226342 + ] + }, + "frontendVersion": "1.28.7", + "VHS_latentpreview": false, + "VHS_latentpreviewrate": 0, + "VHS_MetadataImage": true, + "VHS_KeepIntermediate": true + }, + "version": 0.4 +} \ No newline at end of file diff --git a/example_workflows/DyPE-Flux-workflow.png b/example_workflows/DyPE-Flux-workflow.png new file mode 100644 index 0000000..be9765a Binary files /dev/null and b/example_workflows/DyPE-Flux-workflow.png differ diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..08ed5ee --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +torch \ No newline at end of file diff --git a/src/patch.py b/src/patch.py new file mode 100644 index 0000000..1c3d2e7 --- /dev/null +++ b/src/patch.py @@ -0,0 +1,111 @@ +import torch +import torch.nn as nn +import math +import types +from comfy.model_patcher import ModelPatcher +from comfy import model_sampling +from .rope import get_1d_rotary_pos_embed + + +class FluxPosEmbed(nn.Module): + def __init__(self, theta: int, axes_dim: list[int], method: str = 'yarn', dype: bool = True, dype_exponent: float = 2.0): # Add dype_exponent + super().__init__() + self.theta = theta + self.axes_dim = axes_dim + self.method = method + self.dype = dype if method != 'base' else False + self.dype_exponent = dype_exponent + self.current_timestep = 1.0 + self.base_resolution = 1024 + self.base_patches = (self.base_resolution // 8) // 2 + + def set_timestep(self, timestep: float): + self.current_timestep = timestep + + def forward(self, ids: torch.Tensor) -> torch.Tensor: + n_axes = ids.shape[-1] + emb_parts = [] + pos = ids.float() + freqs_dtype = torch.bfloat16 + + for i in range(n_axes): + axis_pos = pos[..., i] + axis_dim = self.axes_dim[i] + + common_kwargs = {'dim': axis_dim, 'pos': axis_pos, 'theta': self.theta, 'repeat_interleave_real': True, 'use_real': True, 'freqs_dtype': freqs_dtype} + + # Pass the exponent to the RoPE function + dype_kwargs = {'dype': self.dype, 'current_timestep': self.current_timestep, 'dype_exponent': self.dype_exponent} + + if i > 0: + max_pos = axis_pos.max().item() + current_patches = int(max_pos + 1) + + if self.method == 'yarn' and current_patches > self.base_patches: + max_pe_len = torch.tensor(current_patches, dtype=freqs_dtype, device=pos.device) + cos, sin = get_1d_rotary_pos_embed(**common_kwargs, yarn=True, max_pe_len=max_pe_len, ori_max_pe_len=self.base_patches, **dype_kwargs) + elif self.method == 'ntk' and current_patches > self.base_patches: + base_ntk_scale = (current_patches / self.base_patches) + cos, sin = get_1d_rotary_pos_embed(**common_kwargs, ntk_factor=base_ntk_scale, **dype_kwargs) + else: + cos, sin = get_1d_rotary_pos_embed(**common_kwargs) + else: + cos, sin = get_1d_rotary_pos_embed(**common_kwargs) + + cos_reshaped = cos.view(*cos.shape[:-1], -1, 2)[..., :1] + sin_reshaped = sin.view(*sin.shape[:-1], -1, 2)[..., :1] + row1 = torch.cat([cos_reshaped, -sin_reshaped], dim=-1) + row2 = torch.cat([sin_reshaped, cos_reshaped], dim=-1) + matrix = torch.stack([row1, row2], dim=-2) + emb_parts.append(matrix) + + emb = torch.cat(emb_parts, dim=-3) + return emb.unsqueeze(1).to(ids.device) + +def apply_dype_to_flux(model: ModelPatcher, width: int, height: int, method: str, enable_dype: bool, dype_exponent: float, base_shift: float, max_shift: float) -> ModelPatcher: + m = model.clone() + + if not hasattr(m.model.model_sampling, "_dype_patched"): + model_sampler = m.model.model_sampling + if isinstance(model_sampler, model_sampling.ModelSamplingFlux): + patch_size = m.model.diffusion_model.patch_size + latent_h, latent_w = height // 8, width // 8 + padded_h, padded_w = math.ceil(latent_h / patch_size) * patch_size, math.ceil(latent_w / patch_size) * patch_size + image_seq_len = (padded_h // patch_size) * (padded_w // patch_size) + base_seq_len, max_seq_len = 256, 4096 + slope = (max_shift - base_shift) / (max_seq_len - base_seq_len) + intercept = base_shift - slope * base_seq_len + dype_shift = image_seq_len * slope + intercept + + def patched_sigma_func(self, timestep): + return model_sampling.flux_time_shift(dype_shift, 1.0, timestep) + + model_sampler.sigma = types.MethodType(patched_sigma_func, model_sampler) + model_sampler._dype_patched = True + + try: + orig_embedder = m.model.diffusion_model.pe_embedder + theta, axes_dim = orig_embedder.theta, orig_embedder.axes_dim + except AttributeError: + raise ValueError("The provided model is not a compatible FLUX model.") + + new_pe_embedder = FluxPosEmbed(theta, axes_dim, method, enable_dype, dype_exponent) + m.add_object_patch("diffusion_model.pe_embedder", new_pe_embedder) + + sigma_max = m.model.model_sampling.sigma_max.item() + + def dype_wrapper_function(model_function, args_dict): + if enable_dype: + timestep_tensor = args_dict.get("timestep") + if timestep_tensor is not None and timestep_tensor.numel() > 0: + current_sigma = timestep_tensor.item() + if sigma_max > 0: + normalized_timestep = min(max(current_sigma / sigma_max, 0.0), 1.0) + new_pe_embedder.set_timestep(normalized_timestep) + + input_x, c = args_dict.get("input"), args_dict.get("c", {}) + return model_function(input_x, args_dict.get("timestep"), **c) + + m.set_model_unet_function_wrapper(dype_wrapper_function) + + return m \ No newline at end of file diff --git a/src/rope.py b/src/rope.py new file mode 100644 index 0000000..814dd73 --- /dev/null +++ b/src/rope.py @@ -0,0 +1,100 @@ +import torch +import numpy as np +import math + +def find_correction_factor(num_rotations, dim, base, max_position_embeddings): + return (dim * math.log(max_position_embeddings/(num_rotations * 2 * math.pi)))/(2 * math.log(base)) + +def find_correction_range(low_ratio, high_ratio, dim, base, ori_max_pe_len): + low = np.floor(find_correction_factor(low_ratio, dim, base, ori_max_pe_len)) + high = np.ceil(find_correction_factor(high_ratio, dim, base, ori_max_pe_len)) + return max(low, 0), min(high, dim-1) + +def linear_ramp_mask(min_val, max_val, dim): + if min_val == max_val: + max_val += 0.001 + + linear_func = (torch.arange(dim, dtype=torch.float32) - min_val) / (max_val - min_val) + ramp_func = torch.clamp(linear_func, 0, 1) + return ramp_func + +def find_newbase_ntk(dim, base, scale): + return base * (scale ** (dim / (dim - 2))) + +def get_1d_rotary_pos_embed( + dim: int, + pos: torch.Tensor, + theta: float = 10000.0, + use_real=False, + linear_factor=1.0, + ntk_factor=1.0, + repeat_interleave_real=True, + freqs_dtype=torch.float32, + yarn=False, + max_pe_len=None, + ori_max_pe_len=64, + dype=False, + current_timestep=1.0, + dype_exponent=2.0, +): + assert dim % 2 == 0 + device = pos.device + + if yarn and max_pe_len is not None and max_pe_len > ori_max_pe_len: + if not isinstance(max_pe_len, torch.Tensor): + max_pe_len = torch.tensor(max_pe_len, dtype=freqs_dtype, device=device) + + scale = torch.clamp_min(max_pe_len / ori_max_pe_len, 1.0) + + beta_0, beta_1 = 1.25, 0.75 + gamma_0, gamma_1 = 16, 2 + + freqs_base = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=device) / dim)) + freqs_linear = 1.0 / torch.einsum('..., f -> ... f', scale, (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=device) / dim))) + + new_base = find_newbase_ntk(dim, theta, scale) + if new_base.dim() > 0: new_base = new_base.view(-1, 1) + freqs_ntk = 1.0 / torch.pow(new_base, (torch.arange(0, dim, 2, dtype=freqs_dtype, device=device) / dim)) + if freqs_ntk.dim() > 1: freqs_ntk = freqs_ntk.squeeze() + + if dype: + beta_0 = beta_0 ** (dype_exponent * (current_timestep ** dype_exponent)) + beta_1 = beta_1 ** (dype_exponent * (current_timestep ** dype_exponent)) + + low, high = find_correction_range(beta_0, beta_1, dim, theta, ori_max_pe_len) + low, high = max(0, low), min(dim // 2, high) + + freqs_mask = (1 - linear_ramp_mask(low, high, dim // 2).to(device).to(freqs_dtype)) + freqs = freqs_linear * (1 - freqs_mask) + freqs_ntk * freqs_mask + + if dype: + gamma_0 = gamma_0 ** (dype_exponent * (current_timestep ** dype_exponent)) + gamma_1 = gamma_1 ** (dype_exponent * (current_timestep ** dype_exponent)) + + low, high = find_correction_range(gamma_0, gamma_1, dim, theta, ori_max_pe_len) + low, high = max(0, low), min(dim // 2, high) + + freqs_mask = (1 - linear_ramp_mask(low, high, dim // 2).to(device).to(freqs_dtype)) + freqs = freqs * (1 - freqs_mask) + freqs_base * freqs_mask + + else: + theta_ntk = theta * ntk_factor + if dype and ntk_factor > 1.0: + theta_ntk = theta * (ntk_factor ** (dype_exponent * (current_timestep ** dype_exponent))) + + freqs = 1.0 / (theta_ntk ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=device) / dim)) / linear_factor + + freqs = torch.einsum("...s,d->...sd", pos, freqs) + + if use_real and repeat_interleave_real: + freqs_cos = freqs.cos().repeat_interleave(2, dim=-1).float() + freqs_sin = freqs.sin().repeat_interleave(2, dim=-1).float() + + if yarn and max_pe_len is not None and max_pe_len > ori_max_pe_len: + mscale = torch.where(scale <= 1., torch.tensor(1.0), 0.1 * torch.log(scale) + 1.0).to(scale) + freqs_cos, freqs_sin = freqs_cos * mscale, freqs_sin * mscale + return freqs_cos, freqs_sin + elif use_real: + return freqs.cos().float(), freqs.sin().float() + else: + return torch.polar(torch.ones_like(freqs), freqs) \ No newline at end of file