Init: New RoPE types, Chroma-specific changes, & more

This commit is contained in:
Clybius
2026-05-13 14:19:36 -05:00
parent c1df31f614
commit 2595c8d76d
5 changed files with 969 additions and 312 deletions
+88 -112
View File
@@ -1,159 +1,135 @@
<a id="readme-top"></a>
<div align="center">
<h1 align="center">ComfyUI-DyPE</h1>
<img src="https://github.com/user-attachments/assets/4f11966b-86f7-4bdb-acd4-ada6135db2f8" alt="ComfyUI-DyPE Banner" width="70%">
<h1 align="center">ComfyUI-Chroma-RoPE</h1>
<p align="center">
A ComfyUI custom node that implements <strong>DyPE (Dynamic Position Extrapolation)</strong>, enabling FLUX-based models to generate ultra-high-resolution images (4K and beyond) with exceptional coherence and detail.
Advanced Rotary Position Embedding (RoPE) modifications for Chroma/FLUX models
<br />
. Enable ultra-high-resolution image generation with YaRN-like modifications.
<br />
<a href="https://github.com/wildminder/ComfyUI-DyPE/issues/new?labels=bug&template=bug-report---.md">Report Bug</a>
·
<a href="https://github.com/wildminder/ComfyUI-DyPE/issues/new?labels=enhancement&template=feature-request---.md">Request Feature</a>
</p>
</div>
<!-- PROJECT SHIELDS -->
<div align="center">
[![Stargazers][stars-shield]][stars-url]
[![Issues][issues-shield]][issues-url]
[![Forks][forks-shield]][forks-url]
</div>
<br>
## About The Project
DyPE is a novel, training-free method that allows pre-trained diffusion transformers like FLUX to generate images at resolutions far beyond their training data, with no additional sampling cost.
**ComfyUI-Chroma-RoPE** is a ComfyUI custom node that implements advanced RoPE (Rotary Position Embedding) modifications for Chroma and FLUX-based diffusion models. It enables generating images at resolutions far beyond the model's native training scale without additional training.
It works by taking advantage of the spectral progression inherent to the diffusion process. By dynamically adjusting the model's positional encodings at each step, DyPE matches their frequency spectrum with the current stage of the generative process—focusing on low-frequency structures early on and resolving high-frequency details in later steps. This prevents the repeating artifacts and structural degradation typically seen when pushing models beyond their native resolution.
### Key Features
<div align="center">
- **YaRN (Yet another RoPE extensioN)**: Extrapolate position encodings to handle longer sequences/resolutions
- **p-RoPE (Proportional RoPE)**: Selectively truncate low-frequency RoPE dimensions for semantic coherence (Kinda not done right)
- **DyPE**: Timestep-dependent frequency modulation that adapts to the diffusion process
- **Multiple Ramp Functions**: Linear, sigmoid, power-2, and square root blending options
- **Timestep Modulation**: Dynamic theta scaling based on diffusion noise levels
- **Zero Training Required**: Works with pre-trained models out of the box
<img alt="ComfyUI-DyPE example workflow" width="70%" src="https://github.com/user-attachments/assets/e5c1d202-b2c4-474b-b52f-9691ab44c47a" />
<p><sub><i>A simple, single-node integration to patch your FLUX model for high-resolution generation.</i></sub></p>
</div>
This implementation is based on concepts from the original [ComfyUI-DyPE](https://github.com/wildminder/ComfyUI-DyPE) repository and extends it with additional RoPE modification methods.
This node provides a seamless, "plug-and-play" integration of DyPE into any FLUX-based workflow.
## Installation
**✨ Key Features:**
* **True High-Resolution Generation:** Push FLUX models to 4096x4096 and beyond while maintaining global coherence and fine detail.
* **Single-Node Integration:** Simply place the `DyPE for FLUX` node after your model loader to patch the model. No complex workflow changes required.
* **Full Compatibility:** Works seamlessly with your existing ComfyUI workflows, samplers, schedulers, and other optimization nodes like Self-Attention or quantization.
* **Fine-Grained Control:** Exposes key DyPE hyperparameters, allowing you to tune the algorithm's strength and behavior for optimal results at different target resolutions.
* **Zero Inference Overhead:** DyPE's adjustments happen on-the-fly with negligible performance impact.
### Via ComfyUI Manager (Recommended)
<div align="center">
<img alt="Node" width="70%" src="https://github.com/user-attachments/assets/3ef232d2-6268-4e3d-8522-b704dade03ac" />
</div>
1. Open ComfyUI Manager in your ComfyUI interface
2. Click "Install Custom Nodes"
3. Search for `ComfyUI-Chroma-RoPE`
4. Click Install
<p align="right">(<a href="#readme-top">back to top</a>)</p>
### Manual Installation
## 🚀 Getting Started
1. Navigate to your `ComfyUI/custom_nodes/` directory:
```bash
cd ComfyUI/custom_nodes/
```
The easiest way to install is via **ComfyUI Manager**. Search for `ComfyUI-DyPE` and click "Install".
2. Clone this repository:
```bash
git clone https://github.com/Clybius/ComfyUI-Chroma-RoPE.git
```
Alternatively, to install manually:
3. Restart ComfyUI
1. **Clone the Repository:**
## Usage
Navigate to your `ComfyUI/custom_nodes/` directory and clone this repository:
```sh
git clone https://github.com/wildminder/ComfyUI-DyPE.git
```
2. **Start/Restart ComfyUI:**
Launch ComfyUI. No further dependency installation is required.
1. **Load your Chroma/FLUX model** using a standard model loader
2. **Add the Chroma RoPE Patch node** (found under `model_patches/unet`)
3. **Connect the model** output from your loader to the `model` input
4. **Configure parameters** based on your target resolution
5. **Connect to KSampler** and generate
<p align="right">(<a href="#readme-top">back to top</a>)</p>
### Node Parameters
## 🛠️ Usage
| Parameter | Type | Default | Description |
|-----------|------|---------|-------------|
| `model` | Model | Required | The Chroma/FLUX model to patch |
| `method` | Combo | `yarn_freq_stretch` | Position encoding method: `yarn`, `yarn_freq_stretch`, `yarn+dynamic_ntk`, `dynamic_ntk`, `ntk`, `base` |
| `rope_percentage` | Float | 1.0 | p-RoPE proportion (0.0-1.0). 1.0=full RoPE, 0.0=NoPE (semantic only) |
| `dype` | Boolean | True | Enable DyPE (Dynamic Position Extrapolation) with timestep modulation |
| `max_pe_length` | Int | 64 | Original trained position encoding length |
| `yarn_ramp_type` | Combo | `sqrt` | Frequency blending function: `linear`, `sigmoid`, `pow2`, `sqrt` |
| `yarn_ratio` | Float | 1.0 | YaRN scaling ratio multiplier |
| `yarn_beta_fast` | Int | 32 | High-frequency rotation cutoff |
| `yarn_beta_slow` | Int | 2 | Low-frequency rotation cutoff |
| `timestep_modulation` | Boolean | False | Enable timestep-dependent theta scaling |
| `timestep_period_min` | Float | 1000.0 | Theta period at max noise (t=1.0) |
| `timestep_period_max` | Float | 10000.0 | Theta period at min noise (t=0.0) |
| `attn_ratio` | Float | 1.0 | Attention scaling factor for RoPE embeddings |
Using the node is straightforward and designed for minimal workflow disruption.
## Position Encoding Methods
1. **Load Your FLUX Model:** Use a standard `Load Checkpoint` node to load your FLUX model (e.g., `FLUX.1-Krea-dev`).
2. **Add the DyPE Node:** Add the `DyPE for FLUX` node to your graph (found under `model_patches/unet`).
3. **Connect the Model:** Connect the `MODEL` output from your loader to the `model` input of the DyPE node.
4. **Set Resolution:** Set the `width` and `height` on the DyPE node to match the resolution of your `Empty Latent Image`.
5. **Connect to KSampler:** Use the `MODEL` output from the DyPE node as the input for your `KSampler`.
6. **Generate!** That's it. Your workflow is now DyPE-enabled.
### YaRN (yarn)
The default and recommended method. Combines interpolation and extrapolation with frequency-aware blending for smooth high-resolution generation.
> [!NOTE]
> This node specifically patches the **diffusion model (UNet)**. It does not modify the CLIP or VAE models. It is designed exclusively for **FLUX-based** architectures.
### YaRN Frequency Stretch (yarn_freq_stretch)
Novel method that non-linearly stretches the frequency space, preserving high-frequency relationships while extrapolating low frequencies.
### Node Inputs
### YaRN + Dynamic NTK (yarn+dynamic_ntk)
Combines YaRN with Dynamic NTK scaling for aggressive extrapolation scenarios.
* **`model`**: The FLUX model to be patched.
* **`width` / `height`**: The target image resolution. **This must match the resolution set in your `Empty Latent Image` node.**
* **`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.
* **`enable_dype`**: Enables or disables the **dynamic, time-aware** component of DyPE.
* **Enabled (True):** Both the noise schedule and RoPE will be dynamically adjusted throughout sampling. This is the full DyPE algorithm.
* **Disabled (False):** The node will only apply the dynamic noise schedule shift. The RoPE will use a static extrapolation method (e.g., standard YARN). This can be useful for comparison or if you find it works better at certain moderate resolutions.
* **`dype_exponent`**: (λt) Controls the "strength" of the dynamic effect over time. This is the most important tuning parameter.
* `2.0` (Exponential): Recommended for **4K+** resolutions. It's an aggressive schedule that transitions quickly.
* `1.0` (Linear): A good starting point for **~2K-3K** resolutions.
* `0.5` (Sub-linear): A gentler schedule that may work best for resolutions just above the model's native 1K.
* **`base_shift` / `max_shift`** (Advanced): These parameters control the interpolation of the dynamic noise schedule shift (`mu`). The default values (`0.5`, `1.15`) are taken directly from the FLUX architecture and are generally optimal. Adjust only if you are an advanced user experimenting with the noise schedule.
### Dynamic NTK (dynamic_ntk)
Dynamic scaling based on sequence length ratio. More stable for moderate extrapolation.
> [!WARNING]
> It seems the width/height parameters in the node are buggy. Keep the values below 1024x1024; doing so won’t affect your output.
### NTK (ntk)
Standard NTK-aware scaling with fixed extrapolation factor.
<p align="right">(<a href="#readme-top">back to top</a>)</p>
### Base (base)
No extrapolation. Uses original model position encodings.
<p align="center">══════════════════════════════════</p>
## Ramp Functions
Beyond the code, I believe in the power of community and continuous learning. I invite you to join the 'TokenDiff AI News' and 'TokenDiff Community Hub'
The ramp function controls how interpolated and extrapolated frequencies blend:
<table border="0" align="center" cellspacing="10" cellpadding="0">
<tr>
<td align="center" valign="top">
<h4>TokenDiff AI News</h4>
<a href="https://t.me/TokenDiff">
<img width="40%" alt="tokendiff-tg-qw" src="https://github.com/user-attachments/assets/e29f6b3c-52e5-4150-8088-12163a2e1e78" />
</a>
<p><sub>🗞️ AI for every home, creativity for every mind!</sub></p>
</td>
<td align="center" valign="top">
<h4>TokenDiff Community Hub</h4>
<a href="https://t.me/TokenDiff_hub">
<img width="40%" alt="token_hub-tg-qr" src="https://github.com/user-attachments/assets/da544121-5f5b-4e3d-a3ef-02272535929e" />
</a>
<p><sub>💬 questions, help, and thoughtful discussion.</sub> </p>
</td>
</tr>
</table>
- **Linear**: Smooth linear interpolation between regions
- **Sigmoid**: S-curve transition for sharper boundaries
- **Pow2**: Aggressive blending
- **Sqrt**: Conservative blending (gentler transitions)
<p align="center">══════════════════════════════════</p>
## Compatibility
## ⚠️ Known Issues and Limitations
* **FLUX Only:** This implementation is highly specific to the architecture of the FLUX model and will not work on standard U-Net models (like SD 1.5/SDXL) or other Diffusion Transformers.
* **Parameter Tuning:** The optimal `dype_exponent` can vary based on your target resolution. Experimentation is key to finding the best setting for your use case. The default of `2.0` is optimized for 4K.
- **Models**: Chroma, FLUX-based architectures
- **ComfyUI**: Compatible with standard ComfyUI workflows
- **Other Nodes**: Works alongside quantization, attention optimization, and other model patches
<p align="right">(<a href="#readme-top">back to top</a>)</p>
**Not compatible with:** SD 1.5, SDXL, or non-FLUX architectures.
## Known Limitations
- FLUX/Chroma architectures only
- Parameter tuning required for optimal results at different resolutions
- Higher resolutions may require more sampling steps for best quality
## Credits
- Original [ComfyUI-DyPE](https://github.com/wildminder/ComfyUI-DyPE) repository by wildminder
- YaRN paper and implementation concepts
- The ComfyUI team for the extensible platform
<!-- LICENSE -->
## License
The original DyPE project is patent pending. For commercial use or licensing inquiries regarding the underlying method, please contact the [original authors](mailto:noam.issachar@mail.huji.ac.il).
<p align="right">(<a href="#readme-top">back to top</a>)</p>
This project is released under the Apache 2.0 License. See LICENSE file for details.
<!-- ACKNOWLEDGMENTS -->
## Acknowledgments
* **Noam Issachar, Guy Yariv, and the co-authors** for their groundbreaking research and for open-sourcing the [DyPE](https://github.com/guyyariv/DyPE) project.
* **The ComfyUI team** for creating such a powerful and extensible platform for diffusion model research and creativity.
This project builds upon the work from the DyPE research and the original ComfyUI-DyPE implementation. Special thanks to the diffusion model research community for advancing position encoding techniques.
<p align="right">(<a href="#readme-top">back to top</a>)</p>
<!-- MARKDOWN LINKS & IMAGES -->
[stars-shield]: https://img.shields.io/github/stars/wildminder/ComfyUI-DyPE.svg?style=for-the-badge
[stars-url]: https://github.com/wildminder/ComfyUI-DyPE/stargazers
[issues-shield]: https://img.shields.io/github/issues/wildminder/ComfyUI-DyPE.svg?style=for-the-badge
[issues-url]: https://github.com/wildminder/ComfyUI-DyPE/issues
[forks-shield]: https://img.shields.io/github/forks/wildminder/ComfyUI-DyPE.svg?style=for-the-badge
[forks-url]: https://github.com/wildminder/ComfyUI-DyPE/network/members
+139 -44
View File
@@ -2,91 +2,186 @@ import torch
from comfy_api.latest import ComfyExtension, io
from .src.patch import apply_dype_to_flux
class DyPE_FLUX(io.ComfyNode):
class ChromaRoPE(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.
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="DyPE_FLUX",
display_name="DyPE for FLUX",
node_id="ChromaRoPE",
display_name="Chroma RoPE Patch",
category="model_patches/unet",
description="Applies DyPE (Dynamic Position Extrapolation) to a FLUX model for ultra-high-resolution generation.",
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 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."
tooltip="The Chroma model to patch with DyPE.",
),
io.Combo.Input(
"method",
options=["yarn", "ntk", "base"],
default="yarn",
tooltip="Position encoding extrapolation method (YARN recommended).",
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(
"enable_dype",
"dype",
default=True,
label_on="Enabled",
label_off="Disabled",
tooltip="Enable or disable Dynamic Position Extrapolation for RoPE.",
optional=True,
tooltip="Enable Dynamic Position Extrapolation (DyPE) with timestep-dependent frequency modulation. Provides better high-resolution coherence.",
),
io.Float.Input(
"dype_exponent",
default=2.0, min=0.0, max=4.0, step=0.1,
"max_pe_length",
default=64,
min=1,
max=1024,
step=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)."
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(
"base_shift",
default=0.5, min=0.0, max=10.0, step=0.01,
"yarn_ratio",
default=1.00,
min=0.01,
max=10.0,
step=0.01,
optional=True,
tooltip="Advanced: Base shift for the noise schedule (mu). Default is 0.5."
tooltip="YaRN scaling ratio multiplier. Higher values increase extrapolation strength for high-resolution generation.",
),
io.Float.Input(
"max_shift",
default=1.15, min=0.0, max=10.0, step=0.01,
"yarn_beta_fast",
default=32,
min=1,
max=1024,
step=1,
optional=True,
tooltip="Advanced: Max shift for the noise schedule (mu) at high resolutions. Default is 1.15."
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 FLUX model patched with DyPE.",
tooltip="The Chroma 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:
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 FLUX models.")
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, width, height, method, enable_dype, dype_exponent, base_shift, max_shift)
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 DyPEExtension(ComfyExtension):
"""Registers the DyPE node."""
class ChromaRoPEExtension(ComfyExtension):
"""Registers the ChromaRoPE node."""
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [DyPE_FLUX]
return [ChromaRoPE]
async def comfy_entrypoint() -> DyPEExtension:
return DyPEExtension()
async def comfy_entrypoint() -> ChromaRoPEExtension:
return ChromaRoPEExtension()
+5 -5
View File
@@ -1,15 +1,15 @@
[project]
name = "ComfyUI-DyPE"
description = "DyPE for FLUX. Artifact-free 4K+ image generation."
name = "ComfyUI-Chroma-RoPE"
description = "Advanced RoPE modifications for Chroma/FLUX models including DyPE, YaRN, and other RoPE extension methods."
version = "1.0.0"
license = {file = "LICENSE"}
dependencies = ["torch"]
[project.urls]
Repository = "https://github.com/wildminder/ComfyUI-DyPE"
Repository = "https://github.com/Clybius/ComfyUI-Chroma-RoPE"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "wildai"
DisplayName = "ComfyUI-DyPE"
PublisherId = "clybius"
DisplayName = "ComfyUI-Chroma-RoPE"
Icon = ""
+185 -63
View File
@@ -4,98 +4,220 @@ import math
import types
from comfy.model_patcher import ModelPatcher
from comfy import model_sampling
from .rope import get_1d_rotary_pos_embed
from .rope import get_1d_rotary_pos_embed, get_2d_rotary_pos_embed_flexible
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
class ChromaPosEmbed(nn.Module):
"""
A hybrid module for calculating RoPE for multiple positional axes.
It is designed for spatio-temporal data (e.g., T, H, W).
- If initialized for 3 axes, it applies 1D RoPE to the first axis (Time) and
a unified 2D RoPE to the next two axes (Height, Width).
- If initialized for 2 axes, it applies a standard 2D RoPE.
- If initialized for 1 axis, it applies a standard 1D RoPE.
This combines the benefits of independent temporal encoding with principled
2D spatial encoding.
"""
def __init__(
self,
axes_dim: list[int],
theta: float = 10000.0,
method: str = "yarn",
rope_percentage: float = 1.0,
dype: bool = True,
ori_max_pe_len_spatial: int = 64,
yarn_ramp_type: str = "linear",
yarn_ratio: float = 1.0,
yarn_beta_fast: int = 32,
yarn_beta_slow: int = 1,
timestep_modulation: bool = False,
theta_period_min: float = 1.0,
theta_period_max: float = 10000.0,
attn_ratio: float = 1.0,
):
super().__init__()
self.theta = theta
self.n_axes_init = len(axes_dim)
for dim in axes_dim:
assert dim % 2 == 0, (
f"Each dimension in axes_dim must be even, but found {dim}"
)
if self.n_axes_init == 3:
# For 3 axes, the spatial dimension is the sum of the last two.
self.dim_axis0 = axes_dim[0]
self.dim_spatial = axes_dim[1] + axes_dim[2]
assert self.dim_spatial % 2 == 0, (
"Sum of spatial dims (axes 1 and 2) must be even."
)
elif self.n_axes_init == 2:
self.dim_spatial = axes_dim[0] + axes_dim[1]
assert self.dim_spatial % 2 == 0, "Sum of spatial dims must be even."
elif self.n_axes_init == 1:
self.dim_axis0 = axes_dim[0]
else:
raise ValueError(
f"axes_dim must have 1, 2, or 3 elements, but got {self.n_axes_init}"
)
# Store all parameters
self.axes_dim = axes_dim
self.theta = theta
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
self.rope_percentage = rope_percentage
self.dype = dype
self.ori_max_pe_len_spatial = ori_max_pe_len_spatial
self.yarn_ramp_type = yarn_ramp_type
self.yarn_ratio = yarn_ratio
self.yarn_beta_fast = yarn_beta_fast
self.yarn_beta_slow = yarn_beta_slow
self.timestep_modulation = timestep_modulation
self.theta_period_min = theta_period_min
self.theta_period_max = theta_period_max
self.current_timestep = 0.0
self.attn_ratio = attn_ratio
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)
def _get_rotation_matrix(self, cos, sin):
"""Helper to convert cos/sin frequencies to a rotation matrix."""
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)
return torch.stack([row1, row2], dim=-2)
def forward(self, ids: torch.Tensor) -> torch.Tensor:
n_axes_in = ids.shape[-1]
if n_axes_in != self.n_axes_init:
raise ValueError(
f"Input `ids` has {n_axes_in} axes, but module was initialized for {self.n_axes_init} axes."
)
pos = ids.float()
matrix_parts = []
shared_rope_kwargs = {
"theta": self.theta,
"freqs_dtype": torch.float32,
"method": self.method,
"rope_percentage": self.rope_percentage,
"dype": self.dype,
"yarn_ramp_type": self.yarn_ramp_type,
"yarn_ratio": self.yarn_ratio,
"yarn_beta_fast": self.yarn_beta_fast,
"yarn_beta_slow": self.yarn_beta_slow,
"current_timestep": self.current_timestep,
"timestep_modulation": self.timestep_modulation,
"theta_period_min": self.theta_period_min,
"theta_period_max": self.theta_period_max,
"attn_ratio": self.attn_ratio,
}
# --- Hybrid Logic ---
if self.n_axes_init == 3:
# Case 1: 1D (Time) + 2D (Spatial)
# --- Axis 0 (Time): 1D RoPE ---
pos_t = pos[..., 0]
axis0_kwargs = {**shared_rope_kwargs, "dim": self.dim_axis0, "pos": pos_t}
# Time axis is typically not scaled
axis0_kwargs["method"] = "base" if self.method != "base" else "base"
cos_t, sin_t = get_1d_rotary_pos_embed(**axis0_kwargs)
matrix_parts.append(self._get_rotation_matrix(cos_t, sin_t))
# --- Axes 1 & 2 (Height, Width): 2D RoPE ---
pos_x, pos_y = pos[..., 1], pos[..., 2]
spatial_kwargs = {**shared_rope_kwargs, "dim": self.dim_spatial}
# Use one of the spatial axes to determine scaling length
current_max = (pos_y.max().item() + 1 + pos_x.max().item() + 1) // 2
spatial_kwargs["max_pe_len"] = current_max
spatial_kwargs["ori_max_pe_len"] = self.ori_max_pe_len_spatial
cos_spatial, sin_spatial = get_2d_rotary_pos_embed_flexible(
pos_x=pos_x, pos_y=pos_y, **spatial_kwargs
)
matrix_parts.append(self._get_rotation_matrix(cos_spatial, sin_spatial))
elif self.n_axes_init == 2:
# Case 2: Standard 2D RoPE
pos_y, pos_x = pos[..., 0], pos[..., 1]
spatial_kwargs = {**shared_rope_kwargs, "dim": self.dim_spatial}
current_max = (pos_y.max().item() + 1 + pos_x.max().item() + 1) // 2
spatial_kwargs["max_pe_len"] = current_max
spatial_kwargs["ori_max_pe_len"] = self.ori_max_pe_len_spatial
cos_spatial, sin_spatial = get_2d_rotary_pos_embed_flexible(
pos_x=pos_x, pos_y=pos_y, **spatial_kwargs
)
matrix_parts.append(self._get_rotation_matrix(cos_spatial, sin_spatial))
elif self.n_axes_init == 1:
# Case 3: Standard 1D RoPE
pos_t = pos[..., 0]
axis0_kwargs = {**shared_rope_kwargs, "dim": self.dim_axis0, "pos": pos_t}
current_max_len = pos_t.max().item() + 1
axis0_kwargs["max_pe_len"] = current_max_len
axis0_kwargs["ori_max_pe_len"] = (
self.ori_max_pe_len_spatial
) # Assuming spatial scaling applies here too
cos_t, sin_t = get_1d_rotary_pos_embed(**axis0_kwargs)
matrix_parts.append(self._get_rotation_matrix(cos_t, sin_t))
# Concatenate the matrices for different axes along the feature dimension
emb = torch.cat(matrix_parts, dim=-3)
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:
def apply_dype_to_flux(
model: ModelPatcher,
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,
) -> 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)
new_pe_embedder = ChromaPosEmbed(
axes_dim,
theta,
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,
)
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()
+536 -72
View File
@@ -2,99 +2,563 @@ 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)
# Inverse dim formula to find dim based on number of rotations
def find_correction_dim(num_rotations, dim, base=10000, max_position_embeddings=64):
return (dim * math.log(max_position_embeddings / (num_rotations * 2 * math.pi))) / (
2 * math.log(base)
)
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)
# Find dim range bounds based on rotations
def find_correction_range(
low_rot, high_rot, dim, base=10000, max_position_embeddings=64
):
low = math.floor(find_correction_dim(low_rot, dim, base, max_position_embeddings))
high = math.ceil(find_correction_dim(high_rot, dim, base, max_position_embeddings))
return max(low, 0), min(high, dim - 1) # Clamp values just in case
def linear_ramp_mask(min, max, dim):
if min == max:
max += 0.001 # Prevent singularity
linear_func = (torch.arange(dim, dtype=torch.float32) - min) / (max - min)
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 sigmoid_ramp_mask(min_val, max_val, dim):
"""
A sigmoid-based ramp mask.
"""
if min_val == max_val:
max_val += 0.001 # Prevent division by zero
# Scale and shift the linear ramp to be centered around 0 for the sigmoid
linear_func = (torch.arange(dim, dtype=torch.float32) - (min_val + max_val) / 2) / (
max_val - min_val
)
ramp_func = torch.sigmoid(
linear_func * 8
) # The multiplication factor controls the steepness
ramp_func = (ramp_func - ramp_func.min()) / (ramp_func.max() - ramp_func.min())
return ramp_func
def sqrt_ramp_mask(min_val, max_val, dim):
"""
A square root-based ramp mask.
"""
if min_val == max_val:
max_val += 0.001 # Prevent division by zero
linear_func = (torch.arange(dim, dtype=torch.float32) - min_val) / (
max_val - min_val
)
ramp_func = torch.clamp(linear_func, 0, 1).pow(0.5)
return ramp_func
def pow2_ramp_mask(min_val, max_val, dim):
"""
A power of 2-based ramp mask.
"""
if min_val == max_val:
max_val += 0.001 # Prevent division by zero
linear_func = (torch.arange(dim, dtype=torch.float32) - min_val) / (
max_val - min_val
)
ramp_func = torch.clamp(linear_func, 0, 1).pow(2)
return ramp_func
def get_mscale(scale=1):
if scale <= 1:
return 1.0
return 0.1 * math.log(scale) + 1.0
def calculate_base_frequencies(dim, theta, device, dtype):
"""Calculates the base RoPE frequencies."""
return 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=dtype, device=device) / dim))
def calculate_yarn_frequencies_v2(
dim,
max_pe_len,
ori_max_pe_len,
theta,
beta_fast,
beta_slow,
yarn_ratio,
device,
dtype,
dynamic_ntk=False,
ramp="linear",
timestep=None,
):
"""Calculates YaRN-scaled frequencies."""
scale = yarn_ratio * max_pe_len / ori_max_pe_len
beta_0, beta_1 = 1.25, 0.75
gamma_0, gamma_1 = beta_fast, beta_slow
if dynamic_ntk:
# For Dynamic NTK, the scaling factor is based on the ratio of sequence lengths
# and is applied directly to the base.
if scale <= 1:
alpha = 1.0
else:
alpha = ((scale * math.log(ori_max_pe_len)) / math.log(max_pe_len)) ** 2
modified_theta = theta * alpha
else:
modified_theta = theta
# Three RoPE extrapolation/interpolation methods
inv_freq_base = 1.0 / (theta ** (torch.arange(0, dim, 2).float().to(device) / dim))
inv_freq_ntk = 1.0 / (
modified_theta ** (torch.arange(0, dim, 2).float().to(device) / dim)
)
inv_freq_scaled_ntk = 1.0 / (
scale * (modified_theta ** (torch.arange(0, dim, 2).float().to(device) / dim))
)
# inv_freq_scaled = _compute_inv_freq(dim, modified_theta, max_pe_len, ori_max_pe_len, beta_fast, beta_slow).to(device=device, dtype=dtype)
# low, high = find_correction_range(beta_fast, beta_slow, dim, theta, ori_max_pe_len)
beta_0 = beta_0 ** (2.0 * (timestep**2.0)) if timestep is not None else beta_0**2
beta_1 = beta_1 ** (2.0 * (timestep**2.0)) if timestep is not None else beta_1**2
low, high = find_correction_range(beta_0, beta_1, dim, theta, ori_max_pe_len)
low = max(0, low)
high = min(dim // 2, high)
match ramp:
case "linear":
ramp_func = linear_ramp_mask(low, high, dim // 2)
case "sigmoid":
ramp_func = sigmoid_ramp_mask(low, high, dim // 2)
case "pow2":
ramp_func = pow2_ramp_mask(low, high, dim // 2)
case "sqrt":
ramp_func = sqrt_ramp_mask(low, high, dim // 2)
case _:
ramp_func = linear_ramp_mask(low, high, dim // 2)
inv_freq_mask = 1 - ramp_func.to(device=device, dtype=dtype)
freqs = inv_freq_scaled_ntk * (1 - inv_freq_mask) + inv_freq_ntk * inv_freq_mask
gamma_0 = gamma_0 ** (2.0 * (timestep**2.0)) if timestep is not None else gamma_0**2
gamma_1 = gamma_1 ** (2.0 * (timestep**2.0)) if timestep is not None else gamma_1**2
low, high = find_correction_range(gamma_0, gamma_1, dim, theta, ori_max_pe_len)
low = max(0, low)
high = min(dim // 2, high)
match ramp:
case "linear":
ramp_func = linear_ramp_mask(low, high, dim // 2)
case "sigmoid":
ramp_func = sigmoid_ramp_mask(low, high, dim // 2)
case "pow2":
ramp_func = pow2_ramp_mask(low, high, dim // 2)
case "sqrt":
ramp_func = sqrt_ramp_mask(low, high, dim // 2)
case _:
ramp_func = linear_ramp_mask(low, high, dim // 2)
inv_freq_mask = 1 - ramp_func.to(device=device, dtype=dtype)
# print(inv_freq_mask)
final_freqs = freqs * (1 - inv_freq_mask) + inv_freq_base * inv_freq_mask
# final_freqs = inv_freq_base * _get_yarn_scaling_factor(dim, scale, inv_freq_base)
# print(final_freqs)
return final_freqs
def calculate_yarn_freq_stretch(
dim,
max_pe_len,
ori_max_pe_len,
theta,
beta_fast,
beta_slow,
yarn_ratio,
device,
dtype,
ramp="linear",
timestep=None,
stretch_power=3.0,
):
"""
Novel method: "YaRN Frequency Stretch"
Stretches the frequency space non-linearly, keeping the highest
frequencies and their relative phase distances (local frequencies)
untouched/less interpolated.
"""
scale = yarn_ratio * max_pe_len / ori_max_pe_len
if scale <= 1.0:
return 1.0 / (
theta ** (torch.arange(0, dim, 2, dtype=dtype, device=device).float() / dim)
)
# Normalized index n in [0, 1]
n = torch.arange(0, dim, 2, dtype=dtype, device=device).float() / dim
# Base frequencies
inv_freq_base = 1.0 / (theta**n)
# Apply continuous non-linear stretch.
# f'(n) = theta^(-n) * scale^(-n^p)
# At n=0 (high freq), multiplier is 1. Local derivative wrt n is ln(theta), same as base.
# At n=1 (low freq), multiplier is 1/scale.
# Modulate stretch_power depending on the YaRN betas to give user control
actual_power = max(1.0, stretch_power + (beta_fast / 32.0) - 1.0)
# We add a micro-modulation to "leave high frequency components relative to the local frequency untouched"
# By doing a stepped-power function, we protect the relative ratios within octaves.
micro_modulation = (
(torch.sin(n * math.pi * dim / 4.0) / (math.pi * dim / 4.0 + 1e-6))
* 0.1
* (1 - n)
)
g_n = torch.clamp((n**actual_power) - micro_modulation, 0.0, 1.0)
inv_freq_stretched = (1.0 / ((theta * scale) ** n)) * (scale ** (-g_n))
# Blend smoothly with base using the explicit YaRN mask
beta_0, beta_1 = 1.25, 0.75
if timestep is not None:
beta_0 = beta_0 ** (2.0 * (timestep**2.0))
beta_1 = beta_1 ** (2.0 * (timestep**2.0))
low, high = find_correction_range(beta_0, beta_1, dim, theta, ori_max_pe_len)
low = max(0, low)
high = min(dim // 2, high)
match ramp:
case "linear":
ramp_func = linear_ramp_mask(low, high, dim // 2)
case "sigmoid":
ramp_func = sigmoid_ramp_mask(low, high, dim // 2)
case "pow2":
ramp_func = pow2_ramp_mask(low, high, dim // 2)
case "sqrt":
ramp_func = sqrt_ramp_mask(low, high, dim // 2)
case _:
ramp_func = linear_ramp_mask(low, high, dim // 2)
inv_freq_mask = 1 - ramp_func.to(device=device, dtype=dtype)
# Final blend
final_freqs = (
inv_freq_stretched * (1 - inv_freq_mask) + inv_freq_base * inv_freq_mask
)
return final_freqs
def calculate_ntk_frequencies(
dim, max_pe_len, ori_max_pe_len, theta, device, dtype, dynamic=True
):
"""Calculates NTK-scaled or Dynamic-NTK-scaled frequencies."""
scale = max_pe_len / ori_max_pe_len
if dynamic:
# For Dynamic NTK, the scaling factor is based on the ratio of sequence lengths
# and is applied directly to the base.
if scale <= 1:
alpha = 1.0
else:
alpha = ((scale * math.log(ori_max_pe_len)) / math.log(max_pe_len)) ** 2
modified_theta = theta * alpha
else:
# For standard NTK-aware scaling, we scale the base by a fixed factor.
# This factor is often set to the scale itself.
alpha = scale
# The formula from the paper is theta * alpha^(d / (d-2))
modified_theta = theta * (alpha ** (dim / (dim - 2)))
return calculate_base_frequencies(dim, modified_theta, device, dtype)
def find_newbase_ntk(dim, base=10000, scale=1):
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,
# -- Control which method to use ---
method: str = "yarn", # Can be 'yarn', 'yarn_freq_stretch', 'yarn+dynamic_ntk', 'ntk', 'dynamic_ntk', 'base'
rope_percentage: float = 1.0, # NEW: p-RoPE proportion
dype: bool = True, # Dynamic Position Extrapolation
# -- Scaling args ---
max_pe_len: int = None, # Target length, e.g., 256 for a 256x256 latent
ori_max_pe_len: int = 64, # Original trained length, e.g., 64 for a 64x64 latent
# -- YaRN args ---
yarn_ramp_type: str = "linear",
yarn_ratio: float = 1.0,
yarn_beta_fast: int = 32,
yarn_beta_slow: int = 1,
# -- Diffusion Timestep args ---
current_timestep: float = 1.0, # Expected to be in [0, 1]
timestep_modulation: bool = False,
theta_period_min: float = 1000.0, # At t=1 (max noise), theta scale
theta_period_max: float = 10000.0, # At t=0 (no noise), theta scale
# -- Attn scaling --
attn_ratio: float = 1.0,
):
"""
Generates 1D rotary positional embeddings with multiple scaling strategies.
Returns: Tuple of (cos_frequencies, sin_frequencies)
"""
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
# MODIFICATION 2: Timestep-Dependent Frequencies
if timestep_modulation:
# Interpolate theta logarithmically based on the current timestep
# At t=1 (max noise), use a smaller period (higher frequency)
# At t=0 (no noise), use the standard large period
log_min = math.log(theta_period_min)
log_max = math.log(theta_period_max)
log_theta = log_min * current_timestep + log_max * (1.0 - current_timestep)
modified_theta = math.exp(log_theta)
else:
theta_ntk = theta * ntk_factor
if dype and ntk_factor > 1.0:
theta_ntk = theta * (ntk_factor ** (dype_exponent * (current_timestep ** dype_exponent)))
modified_theta = theta
freqs = 1.0 / (theta_ntk ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=device) / dim)) / linear_factor
attn_scale = 1.0
scale = max_pe_len / ori_max_pe_len if max_pe_len is not None else 1.0
# --- Select Frequency Calculation Method ---
# YaRN
if method == "yarn" and max_pe_len is not None:
freqs = calculate_yarn_frequencies_v2(
dim,
max_pe_len,
ori_max_pe_len,
modified_theta,
yarn_beta_fast,
yarn_beta_slow,
yarn_ratio,
device=device,
dtype=freqs_dtype,
ramp=yarn_ramp_type,
timestep=current_timestep if dype else None,
)
attn_scale = get_mscale(scale)
# YaRN Frequency Stretch
elif method == "yarn_freq_stretch" and max_pe_len is not None:
freqs = calculate_yarn_freq_stretch(
dim,
max_pe_len,
ori_max_pe_len,
modified_theta,
yarn_beta_fast,
yarn_beta_slow,
yarn_ratio,
device=device,
dtype=freqs_dtype,
ramp=yarn_ramp_type,
timestep=current_timestep if dype else None,
)
attn_scale = get_mscale(scale)
elif method == "yarn+dynamic_ntk" and max_pe_len is not None:
freqs = calculate_yarn_frequencies_v2(
dim,
max_pe_len,
ori_max_pe_len,
modified_theta,
yarn_beta_fast,
yarn_beta_slow,
yarn_ratio,
device=device,
dtype=freqs_dtype,
dynamic_ntk=True,
ramp=yarn_ramp_type,
timestep=current_timestep if dype else None,
)
attn_scale = get_mscale(scale)
# Dynamic NTK
elif method == "dynamic_ntk" and max_pe_len is not None:
freqs = calculate_ntk_frequencies(
dim,
max_pe_len,
ori_max_pe_len,
modified_theta,
device=device,
dtype=freqs_dtype,
dynamic=True,
)
attn_scale = 1.0 # Dynamic NTK is often more stable
# NTK
elif method == "ntk" and max_pe_len is not None:
freqs = calculate_ntk_frequencies(
dim,
max_pe_len,
ori_max_pe_len,
modified_theta,
device=device,
dtype=freqs_dtype,
dynamic=False,
)
attn_scale = 1.0 # NTK doesn't have a special m-scale, but can be stabilized
# Base
else:
freqs = calculate_base_frequencies(dim, modified_theta, device, freqs_dtype)
# Apply p-RoPE truncation to the calculated frequencies
# This works on top of ANY method (YaRN, NTK, base, etc.)
if rope_percentage < 1.0:
freqs = apply_proportional_rope(freqs, rope_percentage)
# Apply attention scaling
if attn_ratio != 1.0:
attn_scale *= attn_ratio
# Calculate outer product of positions and frequencies
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()
# Handle infinite frequencies (p-RoPE NoPE dimensions)
# Zero out angles where frequency is infinity to prevent NaN in cos/sin
inf_mask = torch.isinf(freqs)
if inf_mask.any():
freqs = torch.where(
inf_mask,
torch.zeros_like(freqs), # Zero angle → cos=1, sin=0 (identity rotation)
freqs,
)
# Return cos and sin components, with interleaved dimensions
freqs_cos = freqs.cos().repeat_interleave(2, dim=-1).to(dtype=torch.float32)
freqs_sin = freqs.sin().repeat_interleave(2, dim=-1).to(dtype=torch.float32)
# Defensive: Check for NaN and replace with identity rotation values
if torch.isnan(freqs_cos).any() or torch.isnan(freqs_sin).any():
freqs_cos = torch.nan_to_num(freqs_cos, nan=1.0) # Identity: cos=1
freqs_sin = torch.nan_to_num(freqs_sin, nan=0.0) # Identity: sin=0
# Apply attention scaling factor directly to the embeddings
# This is equivalent to scaling q and k before the dot product
if attn_scale != 1.0:
freqs_cos *= attn_scale**0.5
freqs_sin *= attn_scale**0.5
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()
# -- 2D Positional Embedding implementation --
def get_2d_rotary_pos_embed(
dim: int, pos_x: torch.Tensor, pos_y: torch.Tensor, **kwargs
):
"""
Generates 2D rotary positional embeddings by splitting the embedding
dimension between the X and Y axes.
Args:
dim (int): Total embedding dimension. Must be divisible by 4.
pos_x (torch.Tensor): Tensor of X positions.
pos_y (torch.Tensor): Tensor of Y positions.
**kwargs: All other arguments are passed to get_1d_rotary_pos_embed.
Returns: Tuple of (cos_frequencies, sin_frequencies)
"""
assert dim % 4 == 0, "Dimension must be divisible by 4 for 2D RoPE"
dim_half = dim // 2
# Get 1D embeddings for X dimension
cos_x, sin_x = get_1d_rotary_pos_embed(dim=dim_half, pos=pos_x, **kwargs)
# Get 1D embeddings for Y dimension
cos_y, sin_y = get_1d_rotary_pos_embed(dim=dim_half, pos=pos_y, **kwargs)
# Concatenate the results
# The first half of the feature dimension is for X, the second half is for Y
final_cos = torch.cat([cos_x, cos_y], dim=-1)
final_sin = torch.cat([sin_x, sin_y], dim=-1)
return final_cos, final_sin
def get_2d_rotary_pos_embed_flexible(
dim: int, pos_x: torch.Tensor, pos_y: torch.Tensor, **kwargs
):
"""
A flexible wrapper for 2D RoPE that handles dimensions not divisible by 4.
It uses a pad-and-trim strategy.
"""
if dim % 4 == 0:
# If dim is perfectly divisible by 4, no changes needed.
return get_2d_rotary_pos_embed(dim, pos_x, pos_y, **kwargs)
else:
return torch.polar(torch.ones_like(freqs), freqs)
# If dim % 4 is not 0 (but must be 2, since we assume dim % 2 == 0)
# 1. Pad the dimension to the next multiple of 4
padded_dim = dim + 2
# 2. Get the positional embeddings for the padded dimension
cos_padded, sin_padded = get_2d_rotary_pos_embed(
padded_dim, pos_x, pos_y, **kwargs
)
# 3. Trim the results back to the original dimension
cos_trimmed = cos_padded[..., :dim]
sin_trimmed = sin_padded[..., :dim]
return cos_trimmed, sin_trimmed
def apply_proportional_rope(
freqs: torch.Tensor, rope_percentage: float = 1.0
) -> torch.Tensor:
"""
Apply p-RoPE truncation to frequency tensor.
p-RoPE keeps only the top 'rope_percentage' of highest frequencies,
setting the rest to infinity (which results in identity rotation / NoPE).
Args:
freqs: Frequency tensor of shape [dim//2] from any RoPE method
rope_percentage: Fraction of dimensions to keep (0.0-1.0)
- 1.0: Keep all (standard RoPE)
- 0.75: Keep top 75%, truncate lowest 25%
- 0.0: Truncate all (NoPE / identity)
Returns:
Modified frequency tensor with infinity padding for NoPE dimensions
"""
if rope_percentage >= 1.0:
# Standard RoPE, no truncation needed
return freqs
if rope_percentage <= 0.0:
# Pure NoPE, all frequencies become infinity
return torch.full_like(freqs, float("inf"))
# Calculate how many frequencies to keep
num_freqs = len(freqs)
keep_count = int(rope_percentage * num_freqs)
if keep_count >= num_freqs:
return freqs
if keep_count <= 0:
return torch.full_like(freqs, float("inf"))
# Create new frequency tensor
# Keep the first 'keep_count' frequencies (which are highest)
# Set remaining to infinity for NoPE
result = freqs.clone()
result[keep_count:] = float("inf")
return result