Эпоха 6 · Генеративка и системы · 2023

57 Speculative Decoding

Fast Inference from Transformers via Speculative Decoding · Leviathan, Kalman, Matias · Google · ICML · и Accelerating LLM Decoding with Speculative Sampling · Chen и др. · DeepMind · 2023
🟧 оригинал выборочно~1.5 чоригинал ↗
Суть за 20 секунд. Маленькая draft-модель дёшево предлагает K токенов подряд; большая target-модель проверяет их все за ОДИН параллельный проход. Каждый draft-токен принимается с вероятностью min(1, p/q); при первом отказе — ресэмпл из «остатка». Итог — вывод ровно из распределения target (lossless), но последовательных проходов большой модели меньше → ускорение ×2–3.

Контекст

Авторегрессия декодирует по одному токену, и каждый — это полный проход большой модели. При этом декодинг memory-bound: время уходит на загрузку весов, а не на арифметику, поэтому проверить сразу K токенов стоит почти как один. Здесь и прячется резерв.

Идея и механизм

Пусть q — дешёвая draft-модель (маленькая или дистиллят), p — целевая. Draft-фаза: q авторегрессионно генерит K кандидатов. Verify-фаза: p за один проход считает свои вероятности для всех K позиций. Приёмка идёт по rejection sampling: draft-токен x принимается с вероятностью min(1, p(x)/q(x)); на первом отказе токен пересэмплируют из нормированного остатка max(0, p − q), а хвост черновика отбрасывают. За один проход target продвигаемся сразу на несколько токенов.

теорвер · алгоритмы Приёмка и почему распределение НЕ портится

Draft предложил токен x ∼ q. Принимаем его с вероятностью

\[ \alpha(x) = \min\!\left(1,\ \frac{p(x)}{q(x)}\right) \]

Если отвергли — берём токен из остаточного распределения (то, что target хочет, но draft недодал):

\[ x \sim \mathrm{norm}\big(\max(0,\ p(\cdot) - q(\cdot))\big) \]

Магия в том, что суммарная вероятность выдать x совпадает с p(x) точно: «принято из q» плюс «доресэмплено из остатка» складываются в целевое распределение —

\[ \Pr[\text{вернуть } x] \;=\; q(x)\,\alpha(x) \;+\; (\text{масса отказа})\cdot p_{\text{res}}(x) \;=\; p(x) \]

Поэтому speculative decoding lossless: результат неотличим от честного сэмплинга из p. Ускорение — от того, что за один дорогой проход target принимается в среднем несколько токенов; чем ближе q к p, тем выше доля приёмки.

Python Один цикл speculative decoding
def spec_step(draft, target, ctx, K):
    xs, qs = [], []
    for _ in range(K):                       # draft предлагает K токенов
        q = draft.probs(ctx + xs); x = sample(q)
        xs.append(x); qs.append(q[x])
    P = target.probs_parallel(ctx, xs)        # target проверяет все K за один проход
    out = []
    for j, x in enumerate(xs):
        if random() < min(1, P[j][x] / qs[j]):
            out.append(x)                     # принят
        else:
            out.append(sample(norm(relu(P[j] - draft_dist[j])))); break  # ресэмпл из остатка, стоп
    return out                                 # распределение = ровно target
draft qдёшево K кандидатов target p1 параллельный проход ✓ приняты (3)✗ → ресэмпл, стоп за один проход target — несколько токенов, распределение то же самое
Draft предлагает K токенов, target проверяет их одним проходом: верный префикс принимается, первую «неугодную» позицию target переписывает из остаточного распределения. Быстрее, но качество — целевой модели.
Аналогия. Стажёр (draft) быстро набрасывает следующие несколько слов черновика. Шеф (target) разом пробегает черновик глазами: совпало с тем, что он и сам бы написал — принимает; первую расходящуюся правку делает сам и на этом останавливает чтение. Итоговый текст — как будто его писал шеф, но времени ушло меньше, потому что проверять пачку быстрее, чем писать по слову.

Почему это важно

Ускорение инференса ×2–3 без потери качества — редкое сочетание; стало стандартным приёмом сёрвинга (vLLM и потомки: Medusa, EAGLE, self-speculation). Урок глубже: раз декодинг memory-bound, параллельная верификация почти бесплатна — и её можно обменять на последовательную генерацию.

Связи

← ускоряет32. Transformer

Атакует ровно авторегрессионную природу #32 — «один проход на токен». Не меняя модель и её распределение, speculative decoding снижает число дорогих последовательных проходов большой модели.

← опирается на28. Knowledge Distillation

Выигрыш тем больше, чем ближе draft к target по распределению — поэтому draft часто делают дистилляцией целевой модели (#28). Хороший «маленький двойник» = высокая доля приёмки = сильное ускорение.

↔ другая ось ускорения48. FlashAttention

Обе ускоряют инференс, но по-разному: FlashAttention удешевляет одно вычисление внимания, speculative decoding сокращает число проходов модели. Ортогональны и складываются в проде.

Вопросы пытливого ума

Почему приёмка сохраняет РОВНО распределение target, а не приближённо?

Это точное следствие rejection sampling: приём с вероятностью min(1, p/q) плюс доресэмпл из остатка max(0, p−q) вместе дают вероятность выдать каждый токен = p(x). Математически это coupling, воспроизводящий целевое распределение, даже когда предложения берутся из q. Поэтому вывод статистически неотличим от честного сэмплинга из target — не «почти», а точно.

Если target всё равно считает все K токенов — откуда ускорение?

Потому что декодинг memory-bound: доминирует загрузка весов модели, а не арифметика. Один проход target по K позициям грузит веса один раз и стоит почти как проход по одной. Приняв в среднем несколько токенов за такой проход, мы сокращаем количество последовательных дорогих шагов — отсюда и выигрыш во времени.

Что если draft-модель плохая?

Тогда доля приёмки мала: target часто отвергает предложения, и выигрыш тает (в пределе — как обычная генерация плюс накладные расходы на draft). Ключ — draft, близкий к target по распределению. Отсюда разные рецепты: дистилляция, self-speculation (та же модель предсказывает вперёд), Medusa-головы, EAGLE — все про то, чтобы предложения чаще принимались.

Что читать в оригинале

Читать ключевое — сам алгоритм приёмки/ресэмпла и доказательство, что распределение сохраняется (это сердце работы). Leviathan и др. (arXiv 2211.17192) дают формальную постановку; Chen и др. (2302.01318) — эквивалентную «speculative sampling» формулировку и результаты на больших моделях.