RoPE Scaling: как растянуть контекст LLM с 4K до 128K без потери качества

opensourceaillmropeattentioncontextit
← Back to Blog

Введение: проблема длинного контекста

Вы скачали 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, который:

  1. Применяет down-scale к частям attention
  2. Добавляет inverse frequency scaling
  3. Позволяет обучить 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. Это меняет правила игры для задач анализа длинных документов.