Files
Clybius-ComfyUI-Chroma-RoPE/__init__.py
T

188 lines
7.1 KiB
Python

import torch
from comfy_api.latest import ComfyExtension, io
from .src.patch import apply_dype_to_flux
class ChromaRoPE(io.ComfyNode):
"""
Applies advanced RoPE (Rotary Position Embedding) modifications to Chroma/FLUX models.
Enables ultra-high-resolution image generation through YaRN, p-RoPE, DyPE, and other
position encoding extrapolation methods.
"""
@classmethod
def define_schema(cls) -> io.Schema:
return io.Schema(
node_id="ChromaRoPE",
display_name="Chroma RoPE Patch",
category="model_patches/unet",
description="Applies YaRN, p-RoPE, DyPE and other RoPE modifications to Chroma/FLUX models for high-resolution generation.",
inputs=[
io.Model.Input(
"model",
tooltip="The Chroma model to patch with DyPE.",
),
io.Combo.Input(
"method",
options=[
"yarn",
"yarn_freq_stretch",
"yarn+dynamic_ntk",
"dynamic_ntk",
"ntk",
"base",
],
default="yarn_freq_stretch",
tooltip="Position encoding extrapolation method (YARN Frequency Stretch recommended).",
),
io.Float.Input(
"rope_percentage",
default=1.0,
min=0.0,
max=1.0,
step=0.01,
tooltip="p-RoPE: Proportion of dimensions to apply RoPE. 1.0=standard RoPE, 0.75=truncate lowest 25%, 0.0=NoPE (pure semantic). Applied on top of selected method.",
),
io.Boolean.Input(
"dype",
default=True,
optional=True,
tooltip="Enable Dynamic Position Extrapolation (DyPE) with timestep-dependent frequency modulation. Provides better high-resolution coherence.",
),
io.Float.Input(
"max_pe_length",
default=64,
min=1,
max=1024,
step=1,
optional=True,
tooltip="Advanced: Max shift for the noise schedule (mu) at high resolutions. Default is 64.",
),
io.Combo.Input(
"yarn_ramp_type",
options=["linear", "sigmoid", "pow2", "sqrt"],
default="sqrt",
tooltip="YaRN ramp function type for frequency blending: linear, sigmoid (smooth transition), pow2 (aggressive), or sqrt (conservative).",
),
io.Float.Input(
"yarn_ratio",
default=1.00,
min=0.01,
max=10.0,
step=0.01,
optional=True,
tooltip="YaRN scaling ratio multiplier. Higher values increase extrapolation strength for high-resolution generation.",
),
io.Float.Input(
"yarn_beta_fast",
default=32,
min=1,
max=1024,
step=1,
optional=True,
tooltip="YaRN beta_fast parameter: rotation cutoff for high-frequency dimensions (32=default, lower=more extrapolation).",
),
io.Float.Input(
"yarn_beta_slow",
default=2,
min=1,
max=1024,
step=1,
optional=True,
tooltip="YaRN beta_slow parameter: rotation cutoff for low-frequency dimensions (2=default, higher=less extrapolation).",
),
io.Boolean.Input(
"timestep_modulation",
default=False,
optional=True,
tooltip="Enable timestep-dependent theta scaling. Modulates frequency periods based on diffusion noise level.",
),
io.Float.Input(
"timestep_period_min",
default=1000.0,
min=1.00,
max=1000000.0,
step=1,
optional=True,
tooltip="Theta period at maximum noise (t=1.0). Lower values = higher frequencies during early denoising.",
),
io.Float.Input(
"timestep_period_max",
default=10000.0,
min=1.00,
max=100000000.0,
step=1,
optional=True,
tooltip="Theta period at minimum noise (t=0.0). Higher values = lower frequencies during final refinement.",
),
io.Float.Input(
"attn_ratio",
default=1.000,
min=0.00,
max=10.0,
step=0.001,
optional=True,
tooltip="Attention scaling factor. Applies attention temperature scaling to RoPE embeddings (1.0=normal).",
),
],
outputs=[
io.Model.Output(
display_name="Patched Model",
tooltip="The Chroma model patched with DyPE.",
),
],
)
@classmethod
def execute(
cls,
model,
method: str = "yarn",
rope_percentage: float = 1.0,
dype: bool = True,
max_pe_length: int = 64,
yarn_ramp_type: str = "linear",
yarn_ratio: float = 1.00,
yarn_beta_fast: int = 32,
yarn_beta_slow: int = 1,
timestep_modulation: bool = False,
timestep_period_min: float = 1000.0,
timestep_period_max: float = 10000.0,
attn_ratio: float = 1.0,
) -> 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 Chroma/FLUX models.")
patched_model = apply_dype_to_flux(
model,
method,
rope_percentage,
dype,
max_pe_length,
yarn_ramp_type,
yarn_ratio,
yarn_beta_fast,
yarn_beta_slow,
timestep_modulation,
timestep_period_min,
timestep_period_max,
attn_ratio,
)
return io.NodeOutput(patched_model)
class ChromaRoPEExtension(ComfyExtension):
"""Registers the ChromaRoPE node."""
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [ChromaRoPE]
async def comfy_entrypoint() -> ChromaRoPEExtension:
return ChromaRoPEExtension()