Init: New RoPE types, Chroma-specific changes, & more
This commit is contained in:
@@ -1,159 +1,135 @@
|
|||||||
<a id="readme-top"></a>
|
<a id="readme-top"></a>
|
||||||
|
|
||||||
<div align="center">
|
<div align="center">
|
||||||
<h1 align="center">ComfyUI-DyPE</h1>
|
<h1 align="center">ComfyUI-Chroma-RoPE</h1>
|
||||||
|
|
||||||
<img src="https://github.com/user-attachments/assets/4f11966b-86f7-4bdb-acd4-ada6135db2f8" alt="ComfyUI-DyPE Banner" width="70%">
|
|
||||||
|
|
||||||
|
|
||||||
<p align="center">
|
<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 />
|
<br />
|
||||||
|
. Enable ultra-high-resolution image generation with YaRN-like modifications.
|
||||||
<br />
|
<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>
|
</p>
|
||||||
</div>
|
</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
|
## 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" />
|
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.
|
||||||
<p><sub><i>A simple, single-node integration to patch your FLUX model for high-resolution generation.</i></sub></p>
|
|
||||||
</div>
|
|
||||||
|
|
||||||
This node provides a seamless, "plug-and-play" integration of DyPE into any FLUX-based workflow.
|
|
||||||
|
|
||||||
**✨ Key Features:**
|
## Installation
|
||||||
* **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.
|
|
||||||
|
|
||||||
<div align="center">
|
### Via ComfyUI Manager (Recommended)
|
||||||
<img alt="Node" width="70%" src="https://github.com/user-attachments/assets/3ef232d2-6268-4e3d-8522-b704dade03ac" />
|
|
||||||
</div>
|
|
||||||
|
|
||||||
<p align="right">(<a href="#readme-top">back to top</a>)</p>
|
1. Open ComfyUI Manager in your ComfyUI interface
|
||||||
|
2. Click "Install Custom Nodes"
|
||||||
|
3. Search for `ComfyUI-Chroma-RoPE`
|
||||||
|
4. Click Install
|
||||||
|
|
||||||
## 🚀 Getting Started
|
### Manual Installation
|
||||||
|
|
||||||
The easiest way to install is via **ComfyUI Manager**. Search for `ComfyUI-DyPE` and click "Install".
|
1. Navigate to your `ComfyUI/custom_nodes/` directory:
|
||||||
|
```bash
|
||||||
|
cd ComfyUI/custom_nodes/
|
||||||
|
```
|
||||||
|
|
||||||
Alternatively, to install manually:
|
2. Clone this repository:
|
||||||
|
```bash
|
||||||
|
git clone https://github.com/Clybius/ComfyUI-Chroma-RoPE.git
|
||||||
|
```
|
||||||
|
|
||||||
1. **Clone the Repository:**
|
3. Restart ComfyUI
|
||||||
|
|
||||||
Navigate to your `ComfyUI/custom_nodes/` directory and clone this repository:
|
## Usage
|
||||||
```sh
|
|
||||||
git clone https://github.com/wildminder/ComfyUI-DyPE.git
|
|
||||||
```
|
|
||||||
2. **Start/Restart ComfyUI:**
|
|
||||||
Launch ComfyUI. No further dependency installation is required.
|
|
||||||
|
|
||||||
<p align="right">(<a href="#readme-top">back to top</a>)</p>
|
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
|
||||||
|
|
||||||
## 🛠️ Usage
|
### Node Parameters
|
||||||
|
|
||||||
Using the node is straightforward and designed for minimal workflow disruption.
|
| 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 |
|
||||||
|
|
||||||
1. **Load Your FLUX Model:** Use a standard `Load Checkpoint` node to load your FLUX model (e.g., `FLUX.1-Krea-dev`).
|
## Position Encoding Methods
|
||||||
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.
|
|
||||||
|
|
||||||
> [!NOTE]
|
### YaRN (yarn)
|
||||||
> 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.
|
The default and recommended method. Combines interpolation and extrapolation with frequency-aware blending for smooth high-resolution generation.
|
||||||
|
|
||||||
### Node Inputs
|
### YaRN Frequency Stretch (yarn_freq_stretch)
|
||||||
|
Novel method that non-linearly stretches the frequency space, preserving high-frequency relationships while extrapolating low frequencies.
|
||||||
|
|
||||||
* **`model`**: The FLUX model to be patched.
|
### YaRN + Dynamic NTK (yarn+dynamic_ntk)
|
||||||
* **`width` / `height`**: The target image resolution. **This must match the resolution set in your `Empty Latent Image` node.**
|
Combines YaRN with Dynamic NTK scaling for aggressive extrapolation scenarios.
|
||||||
* **`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.
|
|
||||||
|
|
||||||
> [!WARNING]
|
### Dynamic NTK (dynamic_ntk)
|
||||||
> It seems the width/height parameters in the node are buggy. Keep the values below 1024x1024; doing so won’t affect your output.
|
Dynamic scaling based on sequence length ratio. More stable for moderate extrapolation.
|
||||||
|
|
||||||
<p align="right">(<a href="#readme-top">back to top</a>)</p>
|
### NTK (ntk)
|
||||||
|
Standard NTK-aware scaling with fixed extrapolation factor.
|
||||||
|
|
||||||
<p align="center">══════════════════════════════════</p>
|
### Base (base)
|
||||||
|
No extrapolation. Uses original model position encodings.
|
||||||
|
|
||||||
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'
|
## Ramp Functions
|
||||||
|
|
||||||
<table border="0" align="center" cellspacing="10" cellpadding="0">
|
The ramp function controls how interpolated and extrapolated frequencies blend:
|
||||||
<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>
|
|
||||||
|
|
||||||
<p align="center">══════════════════════════════════</p>
|
- **Linear**: Smooth linear interpolation between regions
|
||||||
|
- **Sigmoid**: S-curve transition for sharper boundaries
|
||||||
|
- **Pow2**: Aggressive blending
|
||||||
|
- **Sqrt**: Conservative blending (gentler transitions)
|
||||||
|
|
||||||
## ⚠️ Known Issues and Limitations
|
## Compatibility
|
||||||
* **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.
|
|
||||||
|
|
||||||
<p align="right">(<a href="#readme-top">back to top</a>)</p>
|
- **Models**: Chroma, FLUX-based architectures
|
||||||
|
- **ComfyUI**: Compatible with standard ComfyUI workflows
|
||||||
|
- **Other Nodes**: Works alongside quantization, attention optimization, and other model patches
|
||||||
|
|
||||||
|
**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
|
## 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
|
## 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.
|
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.
|
||||||
* **The ComfyUI team** for creating such a powerful and extensible platform for diffusion model research and creativity.
|
|
||||||
|
|
||||||
<p align="right">(<a href="#readme-top">back to top</a>)</p>
|
<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
|
|
||||||
|
|||||||
+140
-45
@@ -2,91 +2,186 @@ import torch
|
|||||||
from comfy_api.latest import ComfyExtension, io
|
from comfy_api.latest import ComfyExtension, io
|
||||||
from .src.patch import apply_dype_to_flux
|
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.
|
Applies advanced RoPE (Rotary Position Embedding) modifications to Chroma/FLUX models.
|
||||||
This allows generating images at resolutions far beyond the model's training scale
|
Enables ultra-high-resolution image generation through YaRN, p-RoPE, DyPE, and other
|
||||||
by dynamically adjusting positional encodings and the noise schedule.
|
position encoding extrapolation methods.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def define_schema(cls) -> io.Schema:
|
def define_schema(cls) -> io.Schema:
|
||||||
return io.Schema(
|
return io.Schema(
|
||||||
node_id="DyPE_FLUX",
|
node_id="ChromaRoPE",
|
||||||
display_name="DyPE for FLUX",
|
display_name="Chroma RoPE Patch",
|
||||||
category="model_patches/unet",
|
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=[
|
inputs=[
|
||||||
io.Model.Input(
|
io.Model.Input(
|
||||||
"model",
|
"model",
|
||||||
tooltip="The FLUX model to patch with DyPE.",
|
tooltip="The Chroma 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(
|
io.Combo.Input(
|
||||||
"method",
|
"method",
|
||||||
options=["yarn", "ntk", "base"],
|
options=[
|
||||||
default="yarn",
|
"yarn",
|
||||||
tooltip="Position encoding extrapolation method (YARN recommended).",
|
"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(
|
io.Boolean.Input(
|
||||||
"enable_dype",
|
"dype",
|
||||||
default=True,
|
default=True,
|
||||||
label_on="Enabled",
|
optional=True,
|
||||||
label_off="Disabled",
|
tooltip="Enable Dynamic Position Extrapolation (DyPE) with timestep-dependent frequency modulation. Provides better high-resolution coherence.",
|
||||||
tooltip="Enable or disable Dynamic Position Extrapolation for RoPE.",
|
|
||||||
),
|
),
|
||||||
io.Float.Input(
|
io.Float.Input(
|
||||||
"dype_exponent",
|
"max_pe_length",
|
||||||
default=2.0, min=0.0, max=4.0, step=0.1,
|
default=64,
|
||||||
|
min=1,
|
||||||
|
max=1024,
|
||||||
|
step=1,
|
||||||
optional=True,
|
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(
|
io.Float.Input(
|
||||||
"base_shift",
|
"yarn_ratio",
|
||||||
default=0.5, min=0.0, max=10.0, step=0.01,
|
default=1.00,
|
||||||
|
min=0.01,
|
||||||
|
max=10.0,
|
||||||
|
step=0.01,
|
||||||
optional=True,
|
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(
|
io.Float.Input(
|
||||||
"max_shift",
|
"yarn_beta_fast",
|
||||||
default=1.15, min=0.0, max=10.0, step=0.01,
|
default=32,
|
||||||
|
min=1,
|
||||||
|
max=1024,
|
||||||
|
step=1,
|
||||||
optional=True,
|
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=[
|
outputs=[
|
||||||
io.Model.Output(
|
io.Model.Output(
|
||||||
display_name="Patched Model",
|
display_name="Patched Model",
|
||||||
tooltip="The FLUX model patched with DyPE.",
|
tooltip="The Chroma model patched with DyPE.",
|
||||||
),
|
),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|
||||||
@classmethod
|
@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.
|
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"):
|
if not hasattr(model.model, "diffusion_model") or not hasattr(
|
||||||
raise ValueError("This node is only compatible with FLUX models.")
|
model.model.diffusion_model, "pe_embedder"
|
||||||
|
):
|
||||||
patched_model = apply_dype_to_flux(model, width, height, method, enable_dype, dype_exponent, base_shift, max_shift)
|
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)
|
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]]:
|
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
@@ -1,15 +1,15 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "ComfyUI-DyPE"
|
name = "ComfyUI-Chroma-RoPE"
|
||||||
description = "DyPE for FLUX. Artifact-free 4K+ image generation."
|
description = "Advanced RoPE modifications for Chroma/FLUX models including DyPE, YaRN, and other RoPE extension methods."
|
||||||
version = "1.0.0"
|
version = "1.0.0"
|
||||||
license = {file = "LICENSE"}
|
license = {file = "LICENSE"}
|
||||||
dependencies = ["torch"]
|
dependencies = ["torch"]
|
||||||
|
|
||||||
[project.urls]
|
[project.urls]
|
||||||
Repository = "https://github.com/wildminder/ComfyUI-DyPE"
|
Repository = "https://github.com/Clybius/ComfyUI-Chroma-RoPE"
|
||||||
# Used by Comfy Registry https://comfyregistry.org
|
# Used by Comfy Registry https://comfyregistry.org
|
||||||
|
|
||||||
[tool.comfy]
|
[tool.comfy]
|
||||||
PublisherId = "wildai"
|
PublisherId = "clybius"
|
||||||
DisplayName = "ComfyUI-DyPE"
|
DisplayName = "ComfyUI-Chroma-RoPE"
|
||||||
Icon = ""
|
Icon = ""
|
||||||
|
|||||||
+193
-71
@@ -4,84 +4,192 @@ import math
|
|||||||
import types
|
import types
|
||||||
from comfy.model_patcher import ModelPatcher
|
from comfy.model_patcher import ModelPatcher
|
||||||
from comfy import model_sampling
|
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):
|
class ChromaPosEmbed(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
|
"""
|
||||||
|
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__()
|
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.axes_dim = axes_dim
|
||||||
|
self.theta = theta
|
||||||
self.method = method
|
self.method = method
|
||||||
self.dype = dype if method != 'base' else False
|
self.rope_percentage = rope_percentage
|
||||||
self.dype_exponent = dype_exponent
|
self.dype = dype
|
||||||
self.current_timestep = 1.0
|
self.ori_max_pe_len_spatial = ori_max_pe_len_spatial
|
||||||
self.base_resolution = 1024
|
self.yarn_ramp_type = yarn_ramp_type
|
||||||
self.base_patches = (self.base_resolution // 8) // 2
|
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):
|
def set_timestep(self, timestep: float):
|
||||||
self.current_timestep = timestep
|
self.current_timestep = timestep
|
||||||
|
|
||||||
|
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)
|
||||||
|
return torch.stack([row1, row2], dim=-2)
|
||||||
|
|
||||||
def forward(self, ids: torch.Tensor) -> torch.Tensor:
|
def forward(self, ids: torch.Tensor) -> torch.Tensor:
|
||||||
n_axes = ids.shape[-1]
|
n_axes_in = ids.shape[-1]
|
||||||
emb_parts = []
|
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()
|
pos = ids.float()
|
||||||
freqs_dtype = torch.bfloat16
|
matrix_parts = []
|
||||||
|
|
||||||
for i in range(n_axes):
|
shared_rope_kwargs = {
|
||||||
axis_pos = pos[..., i]
|
"theta": self.theta,
|
||||||
axis_dim = self.axes_dim[i]
|
"freqs_dtype": torch.float32,
|
||||||
|
"method": self.method,
|
||||||
common_kwargs = {'dim': axis_dim, 'pos': axis_pos, 'theta': self.theta, 'repeat_interleave_real': True, 'use_real': True, 'freqs_dtype': freqs_dtype}
|
"rope_percentage": self.rope_percentage,
|
||||||
|
"dype": self.dype,
|
||||||
# Pass the exponent to the RoPE function
|
"yarn_ramp_type": self.yarn_ramp_type,
|
||||||
dype_kwargs = {'dype': self.dype, 'current_timestep': self.current_timestep, 'dype_exponent': self.dype_exponent}
|
"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,
|
||||||
|
}
|
||||||
|
|
||||||
if i > 0:
|
# --- Hybrid Logic ---
|
||||||
max_pos = axis_pos.max().item()
|
if self.n_axes_init == 3:
|
||||||
current_patches = int(max_pos + 1)
|
# Case 1: 1D (Time) + 2D (Spatial)
|
||||||
|
|
||||||
if self.method == 'yarn' and current_patches > self.base_patches:
|
# --- Axis 0 (Time): 1D RoPE ---
|
||||||
max_pe_len = torch.tensor(current_patches, dtype=freqs_dtype, device=pos.device)
|
pos_t = pos[..., 0]
|
||||||
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)
|
axis0_kwargs = {**shared_rope_kwargs, "dim": self.dim_axis0, "pos": pos_t}
|
||||||
elif self.method == 'ntk' and current_patches > self.base_patches:
|
# Time axis is typically not scaled
|
||||||
base_ntk_scale = (current_patches / self.base_patches)
|
axis0_kwargs["method"] = "base" if self.method != "base" else "base"
|
||||||
cos, sin = get_1d_rotary_pos_embed(**common_kwargs, ntk_factor=base_ntk_scale, **dype_kwargs)
|
cos_t, sin_t = get_1d_rotary_pos_embed(**axis0_kwargs)
|
||||||
else:
|
matrix_parts.append(self._get_rotation_matrix(cos_t, sin_t))
|
||||||
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]
|
# --- Axes 1 & 2 (Height, Width): 2D RoPE ---
|
||||||
sin_reshaped = sin.view(*sin.shape[:-1], -1, 2)[..., :1]
|
pos_x, pos_y = pos[..., 1], pos[..., 2]
|
||||||
row1 = torch.cat([cos_reshaped, -sin_reshaped], dim=-1)
|
spatial_kwargs = {**shared_rope_kwargs, "dim": self.dim_spatial}
|
||||||
row2 = torch.cat([sin_reshaped, cos_reshaped], dim=-1)
|
# Use one of the spatial axes to determine scaling length
|
||||||
matrix = torch.stack([row1, row2], dim=-2)
|
current_max = (pos_y.max().item() + 1 + pos_x.max().item() + 1) // 2
|
||||||
emb_parts.append(matrix)
|
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)
|
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()
|
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:
|
try:
|
||||||
orig_embedder = m.model.diffusion_model.pe_embedder
|
orig_embedder = m.model.diffusion_model.pe_embedder
|
||||||
@@ -89,23 +197,37 @@ def apply_dype_to_flux(model: ModelPatcher, width: int, height: int, method: str
|
|||||||
except AttributeError:
|
except AttributeError:
|
||||||
raise ValueError("The provided model is not a compatible FLUX model.")
|
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)
|
m.add_object_patch("diffusion_model.pe_embedder", new_pe_embedder)
|
||||||
|
|
||||||
sigma_max = m.model.model_sampling.sigma_max.item()
|
sigma_max = m.model.model_sampling.sigma_max.item()
|
||||||
|
|
||||||
def dype_wrapper_function(model_function, args_dict):
|
def dype_wrapper_function(model_function, args_dict):
|
||||||
if enable_dype:
|
timestep_tensor = args_dict.get("timestep")
|
||||||
timestep_tensor = args_dict.get("timestep")
|
if timestep_tensor is not None and timestep_tensor.numel() > 0:
|
||||||
if timestep_tensor is not None and timestep_tensor.numel() > 0:
|
current_sigma = timestep_tensor.item()
|
||||||
current_sigma = timestep_tensor.item()
|
if sigma_max > 0:
|
||||||
if sigma_max > 0:
|
normalized_timestep = min(max(current_sigma / sigma_max, 0.0), 1.0)
|
||||||
normalized_timestep = min(max(current_sigma / sigma_max, 0.0), 1.0)
|
new_pe_embedder.set_timestep(normalized_timestep)
|
||||||
new_pe_embedder.set_timestep(normalized_timestep)
|
|
||||||
|
|
||||||
input_x, c = args_dict.get("input"), args_dict.get("c", {})
|
input_x, c = args_dict.get("input"), args_dict.get("c", {})
|
||||||
return model_function(input_x, args_dict.get("timestep"), **c)
|
return model_function(input_x, args_dict.get("timestep"), **c)
|
||||||
|
|
||||||
m.set_model_unet_function_wrapper(dype_wrapper_function)
|
m.set_model_unet_function_wrapper(dype_wrapper_function)
|
||||||
|
|
||||||
return m
|
return m
|
||||||
|
|||||||
+542
-78
@@ -2,99 +2,563 @@ import torch
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import math
|
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):
|
# Inverse dim formula to find dim based on number of rotations
|
||||||
low = np.floor(find_correction_factor(low_ratio, dim, base, ori_max_pe_len))
|
def find_correction_dim(num_rotations, dim, base=10000, max_position_embeddings=64):
|
||||||
high = np.ceil(find_correction_factor(high_ratio, dim, base, ori_max_pe_len))
|
return (dim * math.log(max_position_embeddings / (num_rotations * 2 * math.pi))) / (
|
||||||
return max(low, 0), min(high, dim-1)
|
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)
|
ramp_func = torch.clamp(linear_func, 0, 1)
|
||||||
return ramp_func
|
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(
|
def get_1d_rotary_pos_embed(
|
||||||
dim: int,
|
dim: int,
|
||||||
pos: torch.Tensor,
|
pos: torch.Tensor,
|
||||||
theta: float = 10000.0,
|
theta: float = 10000.0,
|
||||||
use_real=False,
|
freqs_dtype=torch.float32,
|
||||||
linear_factor=1.0,
|
# -- Control which method to use ---
|
||||||
ntk_factor=1.0,
|
method: str = "yarn", # Can be 'yarn', 'yarn_freq_stretch', 'yarn+dynamic_ntk', 'ntk', 'dynamic_ntk', 'base'
|
||||||
repeat_interleave_real=True,
|
rope_percentage: float = 1.0, # NEW: p-RoPE proportion
|
||||||
freqs_dtype=torch.float32,
|
dype: bool = True, # Dynamic Position Extrapolation
|
||||||
yarn=False,
|
# -- Scaling args ---
|
||||||
max_pe_len=None,
|
max_pe_len: int = None, # Target length, e.g., 256 for a 256x256 latent
|
||||||
ori_max_pe_len=64,
|
ori_max_pe_len: int = 64, # Original trained length, e.g., 64 for a 64x64 latent
|
||||||
dype=False,
|
# -- YaRN args ---
|
||||||
current_timestep=1.0,
|
yarn_ramp_type: str = "linear",
|
||||||
dype_exponent=2.0,
|
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
|
assert dim % 2 == 0
|
||||||
device = pos.device
|
device = pos.device
|
||||||
|
|
||||||
if yarn and max_pe_len is not None and max_pe_len > ori_max_pe_len:
|
# MODIFICATION 2: Timestep-Dependent Frequencies
|
||||||
if not isinstance(max_pe_len, torch.Tensor):
|
if timestep_modulation:
|
||||||
max_pe_len = torch.tensor(max_pe_len, dtype=freqs_dtype, device=device)
|
# Interpolate theta logarithmically based on the current timestep
|
||||||
|
# At t=1 (max noise), use a smaller period (higher frequency)
|
||||||
scale = torch.clamp_min(max_pe_len / ori_max_pe_len, 1.0)
|
# At t=0 (no noise), use the standard large period
|
||||||
|
log_min = math.log(theta_period_min)
|
||||||
beta_0, beta_1 = 1.25, 0.75
|
log_max = math.log(theta_period_max)
|
||||||
gamma_0, gamma_1 = 16, 2
|
log_theta = log_min * current_timestep + log_max * (1.0 - current_timestep)
|
||||||
|
modified_theta = math.exp(log_theta)
|
||||||
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:
|
else:
|
||||||
theta_ntk = theta * ntk_factor
|
modified_theta = theta
|
||||||
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
|
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)
|
freqs = torch.einsum("...s,d->...sd", pos, freqs)
|
||||||
|
|
||||||
if use_real and repeat_interleave_real:
|
# Handle infinite frequencies (p-RoPE NoPE dimensions)
|
||||||
freqs_cos = freqs.cos().repeat_interleave(2, dim=-1).float()
|
# Zero out angles where frequency is infinity to prevent NaN in cos/sin
|
||||||
freqs_sin = freqs.sin().repeat_interleave(2, dim=-1).float()
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
if yarn and max_pe_len is not None and max_pe_len > ori_max_pe_len:
|
# Return cos and sin components, with interleaved dimensions
|
||||||
mscale = torch.where(scale <= 1., torch.tensor(1.0), 0.1 * torch.log(scale) + 1.0).to(scale)
|
freqs_cos = freqs.cos().repeat_interleave(2, dim=-1).to(dtype=torch.float32)
|
||||||
freqs_cos, freqs_sin = freqs_cos * mscale, freqs_sin * mscale
|
freqs_sin = freqs.sin().repeat_interleave(2, dim=-1).to(dtype=torch.float32)
|
||||||
return freqs_cos, freqs_sin
|
|
||||||
elif use_real:
|
# Defensive: Check for NaN and replace with identity rotation values
|
||||||
return freqs.cos().float(), freqs.sin().float()
|
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
|
||||||
|
|
||||||
|
return freqs_cos, freqs_sin
|
||||||
|
|
||||||
|
|
||||||
|
# -- 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:
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user