188 lines
7.1 KiB
Python
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()
|