diff --git a/README.md b/README.md index 7d5e26f..21fe7b1 100644 --- a/README.md +++ b/README.md @@ -1,159 +1,135 @@
- A ComfyUI custom node that implements DyPE (Dynamic Position Extrapolation), 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
+ . Enable ultra-high-resolution image generation with YaRN-like modifications.
- Report Bug
- Β·
- Request Feature
A simple, single-node integration to patch your FLUX model for high-resolution generation.
-ββββββββββββββββββββββββββββββββββ
+### 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 -
- TokenDiff AI News- -ποΈ AI for every home, creativity for every mind! - |
-
- TokenDiff Community Hub- -π¬ questions, help, and thoughtful discussion. - |
-
ββββββββββββββββββββββββββββββββββ
+- **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 -* **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. +## Compatibility - +- **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 -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). - +This project is released under the Apache 2.0 License. See LICENSE file for details. - ## 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. - - - -[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 diff --git a/__init__.py b/__init__.py index 12511c5..ac38e12 100644 --- a/__init__.py +++ b/__init__.py @@ -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.") - - patched_model = apply_dype_to_flux(model, width, height, method, enable_dype, dype_exponent, base_shift, max_shift) + if not hasattr(model.model, "diffusion_model") or not hasattr( + model.model.diffusion_model, "pe_embedder" + ): + raise ValueError("This node is only compatible with Chroma/FLUX models.") + + patched_model = apply_dype_to_flux( + model, + method, + rope_percentage, + dype, + max_pe_length, + yarn_ramp_type, + yarn_ratio, + yarn_beta_fast, + yarn_beta_slow, + timestep_modulation, + timestep_period_min, + timestep_period_max, + attn_ratio, + ) return io.NodeOutput(patched_model) -class 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() \ No newline at end of file + +async def comfy_entrypoint() -> ChromaRoPEExtension: + return ChromaRoPEExtension() diff --git a/pyproject.toml b/pyproject.toml index 01670bc..4ee6d27 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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 = "" diff --git a/src/patch.py b/src/patch.py index 1c3d2e7..144855b 100644 --- a/src/patch.py +++ b/src/patch.py @@ -4,84 +4,192 @@ 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 _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: - n_axes = ids.shape[-1] - emb_parts = [] + 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() - freqs_dtype = torch.bfloat16 + matrix_parts = [] - 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} + 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, + } - if i > 0: - max_pos = axis_pos.max().item() - current_patches = int(max_pos + 1) + # --- Hybrid Logic --- + if self.n_axes_init == 3: + # Case 1: 1D (Time) + 2D (Spatial) - 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) + # --- 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)) - 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) + # --- 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 @@ -89,23 +197,37 @@ def apply_dype_to_flux(model: ModelPatcher, width: int, height: int, method: str 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() - if sigma_max > 0: - normalized_timestep = min(max(current_sigma / sigma_max, 0.0), 1.0) - new_pe_embedder.set_timestep(normalized_timestep) - + timestep_tensor = args_dict.get("timestep") + if timestep_tensor is not None and timestep_tensor.numel() > 0: + current_sigma = timestep_tensor.item() + if sigma_max > 0: + normalized_timestep = min(max(current_sigma / sigma_max, 0.0), 1.0) + new_pe_embedder.set_timestep(normalized_timestep) + input_x, c = args_dict.get("input"), args_dict.get("c", {}) return model_function(input_x, args_dict.get("timestep"), **c) m.set_model_unet_function_wrapper(dype_wrapper_function) - - return m \ No newline at end of file + + return m diff --git a/src/rope.py b/src/rope.py index 814dd73..9ebcd9d 100644 --- a/src/rope.py +++ b/src/rope.py @@ -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, + dim: int, + pos: torch.Tensor, + theta: float = 10000.0, + freqs_dtype=torch.float32, + # -- 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, + ) - 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() + # 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 + + 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: - return torch.polar(torch.ones_like(freqs), freqs) \ No newline at end of file + # 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