This commit is contained in:
WildAi
2025-10-28 00:36:04 +03:00
parent 209c1ed8f3
commit c60e4f1937
8 changed files with 1339 additions and 0 deletions
+28
View File
@@ -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 }}
+163
View File
@@ -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/
+92
View File
@@ -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()
+844
View File
@@ -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
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 849 KiB

+1
View File
@@ -0,0 +1 @@
torch
+111
View File
@@ -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
+100
View File
@@ -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)