[Fix] Preserve YaRN scaling when extending rotary caches (#38786)
This commit is contained in:
@@ -14,6 +14,7 @@ from sglang.kernels.ops.attention.rotary_triton import (
|
||||
from sglang.srt.layers.rotary_embedding.base import RotaryEmbedding
|
||||
from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb
|
||||
from sglang.srt.layers.rotary_embedding.yarn import (
|
||||
_extend_yarn_cache,
|
||||
yarn_find_correction_range,
|
||||
yarn_get_mscale_simple,
|
||||
yarn_linear_ramp_mask,
|
||||
@@ -488,6 +489,14 @@ class YaRNScalingMRotaryEmbedding(MRotaryEmbedding):
|
||||
)
|
||||
return inv_freq
|
||||
|
||||
def _ensure_cos_sin_cache_length(self, needed_max_pos: int):
|
||||
self.cos_sin_cache, _ = _extend_yarn_cache(
|
||||
cache=self.cos_sin_cache,
|
||||
compute_inv_freq=lambda: self._compute_inv_freq(self.scaling_factor),
|
||||
mscale=self.mscale,
|
||||
needed_max_pos=needed_max_pos,
|
||||
)
|
||||
|
||||
def _compute_cos_sin_cache(self) -> torch.Tensor:
|
||||
inv_freq = self._compute_inv_freq(self.scaling_factor)
|
||||
t = torch.arange(
|
||||
|
||||
@@ -18,6 +18,7 @@ from sglang.srt.layers.rotary_embedding.utils import (
|
||||
rotate_neox,
|
||||
)
|
||||
from sglang.srt.layers.rotary_embedding.yarn import (
|
||||
_extend_yarn_cache,
|
||||
yarn_find_correction_range,
|
||||
yarn_get_mscale,
|
||||
yarn_linear_ramp_mask,
|
||||
@@ -419,6 +420,25 @@ class DeepseekScalingRotaryEmbedding(RotaryEmbedding):
|
||||
self.sin_cached_total = torch.sin(emb) * self.mscale
|
||||
return cache
|
||||
|
||||
def _ensure_cos_sin_cache_length(self, needed_max_pos: int):
|
||||
self.cos_sin_cache, rows = _extend_yarn_cache(
|
||||
cache=self.cos_sin_cache,
|
||||
compute_inv_freq=lambda: self._compute_inv_freq(self.scaling_factor),
|
||||
mscale=self.mscale,
|
||||
needed_max_pos=needed_max_pos,
|
||||
)
|
||||
# NPU also consumes full-width cos/sin tables, built before dtype casting.
|
||||
if rows is not None and self.cos_cached_total is not None:
|
||||
cos, sin = rows.chunk(2, dim=-1)
|
||||
self.cos_cached_total = torch.cat(
|
||||
(self.cos_cached_total, cos.repeat(1, 2).to(self.cos_cached_total)),
|
||||
dim=0,
|
||||
)
|
||||
self.sin_cached_total = torch.cat(
|
||||
(self.sin_cached_total, sin.repeat(1, 2).to(self.sin_cached_total)),
|
||||
dim=0,
|
||||
)
|
||||
|
||||
def get_cos_cached_total(self):
|
||||
return self.cos_cached_total
|
||||
|
||||
|
||||
@@ -3,10 +3,11 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Tuple
|
||||
from typing import Callable, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.rotary_embedding.base import RotaryEmbedding
|
||||
|
||||
|
||||
@@ -62,6 +63,31 @@ def yarn_get_mscale(scale: float = 1, mscale: float = 1) -> float:
|
||||
return 0.1 * mscale * math.log(scale) + 1.0
|
||||
|
||||
|
||||
def _extend_yarn_cache(
|
||||
cache: torch.Tensor,
|
||||
compute_inv_freq: Callable[[], torch.Tensor],
|
||||
mscale: float,
|
||||
needed_max_pos: int,
|
||||
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
"""Return the extended cache and uncast rows for auxiliary tables.
|
||||
|
||||
A no-op returns the original cache and None without computing frequencies.
|
||||
The caller supplies YaRN frequencies using its scaling factor and unchanged
|
||||
correction-range bound, rather than the base extension's theta argument.
|
||||
"""
|
||||
if needed_max_pos < cache.shape[0]:
|
||||
return cache, None
|
||||
align = envs.SGLANG_ROPE_CACHE_ALIGN.get()
|
||||
new_len = ((needed_max_pos + align) // align) * align
|
||||
inv_freq = compute_inv_freq().to(cache.device)
|
||||
positions = torch.arange(
|
||||
cache.shape[0], new_len, dtype=inv_freq.dtype, device=cache.device
|
||||
)
|
||||
freqs = torch.einsum("i,j->ij", positions, inv_freq)
|
||||
rows = torch.cat((freqs.cos() * mscale, freqs.sin() * mscale), dim=-1)
|
||||
return torch.cat((cache, rows.to(cache.dtype)), dim=0), rows
|
||||
|
||||
|
||||
class YaRNScalingRotaryEmbedding(RotaryEmbedding):
|
||||
"""RotaryEmbedding extended with YaRN method.
|
||||
|
||||
@@ -134,6 +160,14 @@ class YaRNScalingRotaryEmbedding(RotaryEmbedding):
|
||||
)
|
||||
return inv_freq
|
||||
|
||||
def _ensure_cos_sin_cache_length(self, needed_max_pos: int):
|
||||
self.cos_sin_cache, _ = _extend_yarn_cache(
|
||||
cache=self.cos_sin_cache,
|
||||
compute_inv_freq=lambda: self._compute_inv_freq(self.scaling_factor),
|
||||
mscale=self.mscale,
|
||||
needed_max_pos=needed_max_pos,
|
||||
)
|
||||
|
||||
def _compute_cos_sin_cache(self) -> torch.Tensor:
|
||||
inv_freq = self._compute_inv_freq(self.scaling_factor)
|
||||
t = torch.arange(
|
||||
|
||||
Reference in New Issue
Block a user