RoPE Scaling: как растянуть контекст LLM с 4K до 128K без потери качества
Введение: проблема длинного контекста
Вы скачали Llama 3 8B и обнаружили, что максимальная длина контекста — 8K токенов. Но документ, который нужно проанализировать, содержит 50K токенов. Что делать?
Простое увеличение max_position_embeddings не работает: модель не умеет обрабатывать токены за пределами обученного окна. Паттерны внимания "ломаются" на длинных последовательностях.
Решение — RoPE Scaling — методы, которые позволяют "растянуть" позиционное кодирование на более длинные последовательности.
В этой статье разберём:
- Что такое RoPE (Rotary Positional Embeddings)
- Почему RoPE не работает за пределами обученной длины
- NTK-aware scaling, linear scaling, Yarn
- Как выбрать масштабирование для вашей задачи
- Практические примеры с llama.cpp и vLLM
Позиционное кодирование: зачем оно нужно?
Трансформер не знает порядок
В отличие от RNN, трансформер обрабатывает все токены параллельно. Ему нужно сообщить, где какой токен находится:
"Кот сидит на коврике" ≠ "Коврик сидит на коте"
Без позиций: [Кот, сидит, на, коврике] → bag of words
С позициями: [0, 1, 2, 3] → порядок сохранён
Типы позиционного кодирования
1. Absolute (BERT, Llama 2):
pos_embedding[i] = learned_vector[i]
Проблема: фиксированная длина, нельзя обобщить
2. Sinusoidal (Original Transformer):
pos_embedding[i] = sin/cos(pos / 10000^(i/d_model))
Проблема: плохо работает на длинных последовательностях
3. Rotary (RoPE, Llama 3, Mistral):
pos_embedding[i] = rotation_matrix(pos, dim)
Проблема: тоже ограничен обученной длиной
RoPE (Rotary Positional Embeddings)
Как работает RoPE
RoPE применяет вращение к query и key векторам в зависимости от позиции:
def rope(x, pos, base=10000):
"""
Применяет поворотное кодирование позиции.
x: [batch, seq_len, head_dim] — query или key
pos: [seq_len] — позиции токенов (0, 1, 2, ...)
base: частота для sin/cos
"""
dim = x.shape[-1]
theta = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
# Создаём позиции
positions = pos.unsqueeze(-1) * theta.unsqueeze(0) # [seq_len, dim//2]
cos = positions.cos().unsqueeze(1) # [seq_len, 1, dim//2]
sin = positions.sin().unsqueeze(1)
# Применяем вращение к парам dim
x1 = x[..., 0::2] # чётные
x2 = x[..., 1::2] # нечётные
output = torch.cat([
x1 * cos - x2 * sin,
x1 * sin + x2 * cos
], dim=-1)
return output
Визуализация
Позиция 0: rotation = 0° (нет вращения)
Позиция 1: rotation = θ
Позиция 2: rotation = 2θ
Позиция 100: rotation = 100θ
Позиция 4096: rotation = 4096θ → очень большое вращение
При seq_len > 4096: rotation > 4096θ → паттерны "схлопываются"
Почему RoPE ломается за пределами обученной длины?
Обучение: RoPE видит позиции 0..4095
Инференс: нужно позиции 0..8191
Проблема 1: theta рассчитан для max_len=4096
При pos=8192: rotation = 8192θ → слишком большое вращение
Проблема 2: attention patterns не обучены на длинных расстояниях
Q[i] attention к K[j], где |i-j| > 4096 → не обучено
Проблема 3: perplexity взлетает
Модель "не узнаёт" токены на позиции 5000+
NTK-aware RoPE Scaling
Основная идея
NTK (Nuclear Norm Kernel) scaling — метод, который адаптирует theta для более длинных последовательностей, сохраняя качество attention.
Обычный RoPE:
theta = 10000^(-2/d_model)
rotation(pos) = pos * theta
NTK-aware RoPE:
scale_factor = new_max_len / old_max_len
theta_scaled = theta * (1 + log(scale_factor) / log(base)) ^ (-2/(d_model-2))
rotation(pos) = pos * theta_scaled
Реализация
def rope_ntk_aware(x, pos, base=10000, old_max_len=4096, new_max_len=32768):
"""
NTK-aware RoPE scaling.
Адаптирует theta для поддержки более длинных последовательностей.
"""
dim = x.shape[-1]
# Базовый theta
theta = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
# NTK scaling factor
scale_factor = new_max_len / old_max_len
# Адаптируем theta
if scale_factor > 1:
dim = x.shape[-1]
low_freq_factor = scale_factor ** (dim / (dim - 2))
high_freq_factor = 4 * (low_freq_factor - 1) / (9 * low_freq_factor - 5)
# Для каждой пары dim выбираем theta
for i in range(dim // 2):
freq = 1.0 / theta[i]
if freq < low_freq_factor:
theta[i] = base ** (-2 * i / dim) # original
elif freq > high_freq_factor:
theta[i] = 1.0 / (new_max_len * freq / old_max_len) # scaled
else:
# Interpolation
theta[i] = 1.0 / interpolate(freq, low_freq_factor, high_freq_factor)
# Применяем RoPE с новым theta
positions = pos.unsqueeze(-1) * theta.unsqueeze(0)
cos = positions.cos().unsqueeze(1)
sin = positions.sin().unsqueeze(1)
x1 = x[..., 0::2]
x2 = x[..., 1::2]
return torch.cat([
x1 * cos - x2 * sin,
x1 * sin + x2 * cos
], dim=-1)
llama.cpp реализация
// llama.cpp: NTK scaling для RoPE
static void rope_yarn_corr_dims(
int64_t ndim, int64_t dim,
float mscale, float pitch_scale,
float *thscales
) {
// Вычисляем коррекционные коэффициенты для каждой dim
// mscale = manifest_scale (обычно 1.0)
// pitch_scale = param_alpha (обученный параметр)
}
Yarn (Yet another RoPE scaling)
Основная идея
Yarn — улучшенный метод масштабирования RoPE, который:
- Применяет down-scale к частям attention
- Добавляет inverse frequency scaling
- Позволяет обучить scale factor отдельно
Yarn = RoPE + down-scale + inverse freq scaling
Обычный RoPE:
theta = 10000^(-2/d)
attention(Q, K) = softmax(QK^T / sqrt(d))
Yarn:
theta_scaled = theta / scale
attention(Q, K) = softmax(QK^T / (sqrt(d) * scale))
Параметры Yarn
# Yarn параметры (из config)
yarn_params = {
"original_max_position_embeddings": 4096,
"scale": 8.0, # scale factor
"beta_fast": 32, # fast beta
"beta_slow": 1, # slow beta
"mscale": 1.0, # manifest scale
"mscale_all_dim": 1.0 # scale all dimensions
}
Как подобрать scale?
Метод подбора scale:
1. Начните с original_max_len = 4096
2. Определите target_max_len = 32768
3. scale = target_max_len / original_max_len = 8.0
4. Для llama.cpp:
--rope-scaling yarn --rope-freq-scale 0.125
(0.125 = 1/8)
5. Для vLLM:
scaling_policy = "yarn"
original_max_position_embeddings = 4096
scale = 8
Linear Scaling (Simple RoPE Scaling)
Самый простой метод
Linear Scaling:
theta_new = theta_old * (old_max_len / new_max_len)
Пример:
old_max_len = 4096
new_max_len = 32768
ratio = 4096 / 32768 = 0.125
theta_new = theta_old * 0.125
llama.cpp: rope_freq_scale
# rope_freq_scale = 1.0 / scaling_factor
# scaling_factor = 8 → rope_freq_scale = 0.125
./server -m model.gguf \
--ctx-size 32768 \
--rope-freq-scale 0.125
# В llama.cpp:
float rope_freq_scale = params.rope_freq_scale;
if (rope_freq_scale != 1.0f) {
// Масштабируем theta
theta *= rope_freq_scale;
}
Сравнение методов
Точность на длинных контекстах
Метод | 4K | 8K | 16K | 32K | 64K | 128K
------------------|-----|-----|-----|-----|-----|-----
No scaling | 1.0 | 1.3 | 2.1 | 5.8 | OOM | OOM
Linear scaling | 1.0 | 1.1 | 1.4 | 2.0 | 3.5 | 6.2
NTK-aware | 1.0 | 1.0 | 1.1 | 1.3 | 1.8 | 2.8
Yarn | 1.0 | 1.0 | 1.0 | 1.2 | 1.5 | 2.1
Перplexity ratio (относительно 4K). Чем ближе к 1.0 — тем лучше.
Скорость инференса
Метод | Overhead
------------------|----------
No scaling | 0%
Linear scaling | 0%
NTK-aware | <1%
Yarn | <1%
Все методы имеют negligible overhead — вычисления те же самые.
Практика: настройка в разных инструментах
llama.cpp
# Linear scaling
./server -m model.gguf --ctx-size 32768 --rope-freq-scale 0.125
# Yarn scaling
./server -m model.gguf --ctx-size 32768 --rope-scaling yarn --rope-freq-scale 0.125
# NTK-aware (автоматически при --ctx-size > original)
./server -m model.gguf --ctx-size 32768
# Автоматическое определение
./server -m model.gguf --ctx-size 32768 --rope-scaling dynamic
vLLM
from vllm import LLM, SamplingParams
llm = LLM(
model="meta-llama/Llama-3-8B-Instruct",
max_model_len=32768, # Устанавливаем целевую длину
rope_scaling={
"type": "yarn",
"original_max_position_embeddings": 4096,
"scale": 8
}
)
sampling_params = SamplingParams(max_tokens=1024)
Ollama
# Ollama автоматически применяет scaling при увеличении ctx
ollama run llama3.1 "Прочитай и проанализируй этот документ..."
# Или через Modelfile:
FROM llama3.1
PARAMETER num_ctx 32768
Transformers (HuggingFace)
from transformers import AutoModelForCausalLM, AutoTokenizer
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3-8B-Instruct",
rope_scaling={
"type": "yarn",
"original_max_position_embeddings": 4096,
"scale": 8
}
)
model.config.max_position_embeddings = 32768
Когда какое масштабирование?
Linear Scaling
- Подходит: быстрое прототипирование, small scale (2x-4x)
- Не подходит: large scale (>8x), production
- Плюсы: просто, предсказуемо
- Минусы: качество падает на больших scale
NTK-aware
- Подходит: medium scale (4x-16x), баланс качество/скорость
- Не подходит: extreme scale (>32x)
- Плюсы: лучше linear на medium scale
- Минусы: сложнее в настройке
Yarn
- Подходит: large scale (8x-32x), production
- Не подходит: когда нужен quick hack
- Плюсы: лучшее качество на больших scale
- Минусы: нужно подобрать scale factor
Ограничения масштабирования
Память
KV Cache растёт с ctx-size:
Llama 3 8B, FP16:
ctx=4K: KV cache = 2 × 8 × 4096 × 4096 × 2 = 512 MB
ctx=32K: KV cache = 2 × 8 × 32768 × 4096 × 2 = 4 GB
ctx=128K: KV cache = 2 × 8 × 131072 × 4096 × 2 = 16 GB
При batch_size=1:
ctx=128K → 16 GB VRAM только на KV cache!
Attention complexity
Attention вычисление: O(seq_len^2)
seq_len = 4096: 4096^2 = 16M operations
seq_len = 32768: 32768^2 = 1073M operations (×67!)
seq_len = 131072: 131072^2 = 17B operations (×1024!)
Решения: Flash Attention 2
Flash Attention 2 оптимизирует attention:
- Tiled computation: разбиваем на чанки
- Recomputation: не храним intermediate attention map
- Parallel reduction: быстрее softmax
Результат:
seq_len=32K: ×3-5 быстрее
seq_len=128K: ×5-10 быстрее
Практические рекомендации
Для чтения документов
Задача: прочитать и ответить на вопросы по документу
Рекомендация:
1. Yarn scaling, scale = target_len / original_len
2. ctx-size = len(document) + len(questions) + 512
3. Flash Attention 2 включён
4. Batch size = 1 (KV cache растёт с ctx)
Пример:
Документ: 20K токенов
Вопросы: 512 токенов
ctx-size = 21500 → округляем до 32768
rope-freq-scale = 4096/32768 = 0.125
Для чатов с историей
Задача: длинный чат с сохранением контекста
Рекомендация:
1. NTK-aware scaling (автоматический)
2. ctx-size = len(history) + max_response + 1024
3. Trim history при превышении 80% ctx-size
Пример:
История: 28K токенов
max_response: 2048
ctx-size = 32768
Для генерации кода
Задача: анализ и генерация большого кодового файла
Рекомендация:
1. Linear scaling (code не требует глубокого контекста)
2. ctx-size = len(file) + 1024
3. chunking: разбиваем файл на части по 4K-8K
Пример:
Файл: 10K строк → ~8000 токенов
ctx-size = 16384
rope-freq-scale = 4096/16384 = 0.25
Итоги
- RoPE ограничен обученной длиной (обычно 4K)
- NTK-aware, Yarn, Linear — три основных метода масштабирования
- Yarn даёт лучшее качество на больших scale (8x-32x)
- Linear scaling прост, но падает на scale > 8x
- KV cache и attention complexity растут с ctx-size
- Flash Attention 2 обязателен для ctx > 16K
- Для production: Yarn + Flash Attention 2
- Для quick experiments: Linear scaling достаточно
Правильное масштабирование RoPE позволяет получить 128K контекст из модели, обученной на 4K. Это меняет правила игры для задач анализа длинных документов.