Beam Search: как поиск по нескольким лучам улучшает качество генерации

opensourceaillmdecodingsamplingit
← Back to Blog

Введение: почему Greedy Search — не всегда лучший выбор?

Когда LLM генерирует текст, она делает это токкен за токеном. На каждом шаге модель выдаёт распределение вероятностей по всему vocab (50K+ токенов). Простейший способ выбрать следующий токен — взять тот, у которого вероятность выше всего. Это Greedy Search.

Greedy Search:
  Шаг 1: "Привет" → ["как", "мир", "я", "ты", "все"] → выбираем "как" (0.35)
  Шаг 2: "как" → ["дела", "животные", "погода", "код", "завтрак"] → выбираем "дела" (0.28)
  Шаг 3: "дела" → ["?", "хорошо", "плохо", "так", "отлично"] → выбираем "?" (0.42)
  
  Результат: "Привет, как дела?"

Проблема Greedy Search: он жадный. Не видит, что выбор "я" на шаге 1 может привести к более вероятной последовательности в целом.

Greedy:      "Привет, как дела?" → P = 0.35 × 0.28 × 0.42 = 0.041
Alternative: "Привет, я мир"   → P = 0.15 × 0.60 × 0.50 = 0.045

Greedy выбрал худшую последовательность, потому что не смотрел вперёд.


Проблема: почему Greedy Search плох?

Локальный оптимум ≠ глобальный

Последовательность A (Greedy):
  Токен 1: "Код" (0.40) → локальный максимум
  Токен 2: "на" (0.35)
  Токен 3: "Python" (0.30)
  P(A) = 0.40 × 0.35 × 0.30 = 0.042

Последовательность B (не greedy):
  Токен 1: "Пиши" (0.25) → пропущен greedy
  Токен 2: "код" (0.50)
  Токен 3: "на" (0.45)
  Токен 4: "Python" (0.55)
  P(B) = 0.25 × 0.50 × 0.45 × 0.55 = 0.031
  
  Но если B короче:
  Токен 1: "Пиши" (0.25)
  Токен 2: "на" (0.50)
  Токен 3: "Python" (0.60)
  P(B') = 0.25 × 0.50 × 0.60 = 0.075 > 0.042

Ключевая идея: лучший токен на шаге N может быть частью худшей последовательности в целом. И наоборот — менее вероятный токен на шаге N может открыть путь к очень вероятной последовательности.

Экспоненциальный рост

Vocab size: 50,000
Sequence length: 128

Всего возможных последовательностей: 50,000^128 ≈ 10^632

Это больше, чем атомов во Вселенной (~10^80).

Перебрать все варианты невозможно. Нужен компромисс.

Решение: Beam Search

Основная идея

Вместо одного лучшего токена на каждом шаге — сохраняем K лучших последовательностей (beam width).

Beam Search (K=3):

Шаг 1: Генерируем топ-3 токена для "Привет"
  ["как", "я", "мир"]
  
  Сохраняем 3 луча:
  Beam 1: "Привет, как" (P=0.35)
  Beam 2: "Привет, я"   (P=0.15)
  Beam 3: "Привет, мир" (P=0.10)

Шаг 2: Для каждого луча генерируем топ-3 продолжения
  Beam 1 + ["дела", "животные", "погода"]
  Beam 2 + ["тоже", "хочу", "люблю"]
  Beam 3 + ["красивый", "большой", "новый"]
  
  Всего 9 комбинаций. Выбираем топ-3:
  Beam 1: "Привет, как дела" (P=0.098)
  Beam 2: "Привет, я тоже"   (P=0.090)
  Beam 3: "Привет, я хочу"   (P=0.075)

Шаг 3: Повторяем...

Алгоритм

def beam_search(model, prompt, K=5, max_len=128):
    """
    K: beam width (число лучей)
    """
    # Инициализация: prompt → один луч
    beams = [(prompt, 1.0, [])]  # (text, prob, tokens)
    
    for step in range(max_len):
        candidates = []
        
        # Расширяем каждый луч
        for text, prob, tokens in beams:
            next_token_dist = model.predict(text)  # P(next | tokens)
            
            # Берём топ-K продолжений для каждого луча
            for token, p in next_token_dist.topk(K):
                new_text = text + model.decode(token)
                new_prob = prob * p
                new_tokens = tokens + [token]
                candidates.append((new_text, new_prob, new_tokens))
        
        # Сортируем все кандидаты и берём топ-K
        candidates.sort(key=lambda x: x[1], reverse=True)
        beams = candidates[:K]
        
        # Проверка: если все лучи завершены — стоп
        if all(is_eos(token) for _, _, tokens in beams for token in tokens[-5:]):
            break
    
    # Возвращаем луч с наивысшей вероятностью
    return max(beams, key=lambda x: x[1])

Визуализация

Beam Search (K=3), 4 шага:

Шаг 0: "Привет"
        │
        ├─ Beam 1: "Привет, как" (0.35)
        ├─ Beam 2: "Привет, я"   (0.15)
        └─ Beam 3: "Привет, мир" (0.10)
        
Шаг 1: Каждый луч → 3 продолжения = 9 кандидатов
        │
        ├─ Beam 1: "Привет, как дела" (0.098)
        ├─ Beam 2: "Привет, я тоже"   (0.090)
        └─ Beam 3: "Привет, я хочу"   (0.075)
        
Шаг 2: 9 кандидатов → 27 → топ-3
        │
        ├─ Beam 1: "Привет, как дела?" (0.041)
        ├─ Beam 2: "Привет, я тоже хочу" (0.040)
        └─ Beam 3: "Привет, я люблю код" (0.038)

Beam Search vs Greedy Search: сравнение

Качество генерации

Prompt: "Напиши функцию сортировки"

Greedy (K=1):
  "Напиши функцию сортировки на Python. Вот пример:"
  "def sort(arr):"
  "    for i in range(len(arr)):"
  "        for j in range(i + 1, len(arr)):"
  "            if arr[i] > arr[j]:"
  "                arr[i], arr[j] = arr[j], arr[i]"
  "    return arr"

Beam Search (K=5):
  "Напиши функцию сортировки на Python. Вот пример с использованием алгоритма быстрой сортировки:"
  "def quicksort(arr):"
  "    if len(arr) <= 1:"
  "        return arr"
  "    pivot = arr[len(arr) // 2]"
  "    left = [x for x in arr if x < pivot]"
  "    middle = [x for x in arr if x == pivot]"
  "    right = [x for x in arr if x > pivot]"
  "    return quicksort(left) + middle + quicksort(right)"

Beam Search выбрал более информативный и качественный ответ.

Метрики на BLEU/ROUGE

Модель          | BLEU-4 | ROUGE-1 | ROUGE-L
----------------|--------|---------|--------
Greedy (K=1)    | 12.3   | 35.2    | 32.1
Beam (K=3)      | 14.7   | 37.8    | 34.5
Beam (K=10)     | 15.9   | 39.1    | 35.8
Nucleus (p=0.9) | 13.8   | 36.5    | 33.2
Random          | 8.1    | 28.4    | 25.7

Beam Search consistently улучшает метрики на 15-20% по сравнению с Greedy.


Компромиссы Beam Search

Вычислительная стоимость

Greedy Search (K=1):
  На каждом шаге: 1 forward pass
  Для 128 токенов: 128 forward passes
  
Beam Search (K=5):
  На каждом шаге: 5 forward passes (или 1 с batch=5)
  Для 128 токенов: 640 forward passes
  
Beam Search (K=10):
  На каждом шаге: 10 forward passes
  Для 128 токенов: 1280 forward passes
  
Время инференса:
  Greedy:   ~1.2 сек (128 токенов)
  Beam K=5: ~6.0 сек
  Beam K=10: ~12.0 сек

Проблема: перебор лучей

K=100:
  Качество не улучшается значительно после K=20
  Но стоимость растёт линейно
  
K=1000:
  Практически бесполезно
  99% лучей — избыточны
BLEU vs Beam Width:
  K=1:  12.3
  K=2:  13.5 (+9.8%)
  K=3:  14.7 (+8.9%)
  K=5:  15.9 (+8.2%)
  K=10: 16.4 (+3.1%)
  K=20: 16.7 (+1.8%)
  K=50: 16.9 (+1.2%)
  K=100: 17.0 (+0.6%)
  
  Диминуishing returns после K=10!

Нормализация по длине

Проблема: короткие последовательности имеют преимущество

Последовательность A (10 токенов):
  P = 0.3^10 = 0.0000059
  
Последовательность B (20 токенов):
  P = 0.3^20 = 0.00000000035
  
P(B) в 17 миллионов раз меньше, хотя каждый токен одинаково вероятен!

Решение: log probability + length normalization

def score_beam(prob, length, alpha=0.5):
    """
    Нормализованный скор луча.
    
    alpha: коэффициент нормализации длины
      alpha=0: только log prob (предпочитает короткие)
      alpha=1: средняя log prob (баланс)
      alpha>1: предпочитает длинные
    
    score = (1/length^alpha) * log(prob)
    """
    log_prob = math.log(prob)
    normalized = log_prob / (length ** alpha)
    return normalized

# Практический выбор:
# alpha=0.6 — хороший компромисс для большинства задач
def beam_search_normalized(model, prompt, K=5, max_len=128, alpha=0.6):
    beams = [(prompt, 0.0, [])]  # log prob, инициализируем 0
    
    for step in range(max_len):
        candidates = []
        
        for text, log_prob, tokens in beams:
            next_token_dist = model.predict(text)
            
            for token, p in next_token_dist.topk(K):
                new_log_prob = log_prob + math.log(p)
                length = len(tokens) + 1
                normalized_score = new_log_prob / (length ** alpha)
                
                new_tokens = tokens + [token]
                candidates.append((text, new_log_prob, normalized_score, new_tokens))
        
        # Сортируем по нормализованному скорy
        candidates.sort(key=lambda x: x[2], reverse=True)
        beams = [(t, lp, s, tok) for t, lp, s, tok in candidates[:K]]
    
    return max(beams, key=lambda x: x[1])

Практическое использование

Beam Search в Hugging Face Transformers

from transformers import AutoModelForCausalLM, AutoTokenizer

model = AutoModelForCausalLM.from_pretrained("mistralai/Mixtral-8x7B-Instruct-v0.1")
tokenizer = AutoTokenizer.from_pretrained("mistralai/Mixtral-8x7B-Instruct-v0.1")

inputs = tokenizer("Напиши функцию сортировки", return_tensors="pt")

# Greedy Search
greedy = model.generate(**inputs, do_sample=False, max_length=128)

# Beam Search
beam = model.generate(
    **inputs,
    do_sample=False,
    max_length=128,
    num_beams=5,        # K=5
    early_stopping=True # останавливать когда все лучи завершены
)

# С нормализацией длины
beam_norm = model.generate(
    **inputs,
    do_sample=False,
    num_beams=5,
    length_penalty=0.6,  # alpha
    early_stopping=True
)

Когда использовать Beam Search?

ДА:
  ✓ Machine translation (BLEU важен)
  ✓ Summarization (ROUGE важен)
  ✓ Когда качество важнее скорости
  ✓ Offline generation (не real-time)
  ✓ Evaluation / benchmarking

НЕТ:
  ✗ Chat (нужна разнообразность)
  ✗ Real-time applications (слишком медленно)
  ✗ Creative writing (слишком детерминированно)
  ✗ Streaming (нельзя ждать все лучи)

Beam Search + Sampling

# Комбинация: beam search + top-k sampling
# Сохраняем K лучших лучей, но на каждом шаге выбираем
# токен из топ-K по вероятности

model.generate(
    **inputs,
    num_beams=5,
    do_sample=True,       # добавляем случайность
    top_k=50,             # рассматриваем только топ-50 токенов
    temperature=0.7,      # контролируем случайность
    length_penalty=0.6
)

Альтернативы Beam Search

Contrastive Decoding

# Усиливаем "хорошие" токены, ослабляем "плохие"
# Используем две модели: strong и weak

strong_tokens = strong_model.predict(context)
weak_tokens = weak_model.predict(context)

# Contrastive scores
contrastive_scores = strong_tokens - weak_tokens

# Выбираем по contrastive score
best_token = contrastive_scores.argmax()

Self-Correction Decoding

# Генерируем N вариантов, выбираем лучший
candidates = []
for i in range(10):
    candidate = model.generate(prompt, do_sample=True)
    candidates.append(candidate)

# Оцениваем каждый вариант
scores = [evaluator(candidate) for candidate in candidates]
best = candidates[scores.argmax()]

Speculative Beam Search

# Комбинируем speculative decoding с beam search
# Small model предлагает кандидаты (draft)
# Large model верифицирует (verify)

drafts = small_model.beam_search(prompt, K=5)
verified = large_model.verify(drafts)
best = verified[0]

Итоги

  • Beam Search сохраняет K лучших последовательностей на каждом шаге
  • Улучшает BLEU/ROUGE на 15-20% по сравнению с Greedy
  • Стоимость растёт линейно с K, но качество — логарифмически
  • K=5-10 — хороший компромисс
  • Length normalization (alpha=0.6) предотвращает слишком короткие ответы
  • Отлично для translation/summarization, плохо для chat
  • В production: редко используется из-за скорости, но для offline — стандарт

Beam Search — классический алгоритм декодирования, который остаётся актуальным в 2026 году, особенно в комбинации с sampling и speculative decoding.