Математика глубокого обучения

LoRA

Малый ранг как способ дообучения — и что именно он экономит, а что нет

Шаг 86 из 117 · ~28 мин

Задача

Полное дообучение модели на 7 миллиардов параметров требует держать в памяти не только веса. Adam хранит два момента, плюс сами градиенты — то есть четыре массива размера модели вместо одного. В fp32 это порядка 100 гигабайт под то, что содержательно является небольшой поправкой.

Отсюда вопрос: нельзя ли ограничить поправку так, чтобы она была маленькой по числу параметров, оставаясь при этом полноразмерной по действию?

W=W+αrBA,BRm×r,  ARr×nW' = \htmlData{k=frozen}{W} + \htmlData{k=scale}{\frac{\alpha}{r}}\,\htmlData{k=update}{BA}, \qquad B \in \mathbb{R}^{m\times r},\; A \in \mathbb{R}^{r\times n}

имеет ранг не выше rr и содержит r(n+m)r(n+m) чисел вместо nmnm. Всё остальное в этом уроке — следствия этой замены.

Арифметика ранга

Экономия равна r(n+m)nm\frac{r(n+m)}{nm}, и у неё есть точка безубыточности: ранг, при котором две тонкие матрицы содержат столько же чисел, сколько одна большая.

r=nmn+m,для квадратной n=m:r=n2r^\ast = \frac{nm}{n+m}, \qquad \text{для квадратной } n = m: \quad r^\ast = \frac{n}{2}

Для матрицы 4096×40964096 \times 4096 это 20482048. Практические ранги — от 44 до 6464, то есть в тридцать–пятьсот раз меньше точки безубыточности. Пространства тут столько, что выбор между r=8r = 8 и r=16r = 16 не про память вовсе.

W заморожено+BA
ранг r 8
всего в W
2147.48M
обучаемых
8.39M
доля
0.39%
безубыточный ранг
2048
Adam: LoRA / полное
96 / 24576 МБ
Тридцать два слоя по четыре матрицы 4096×4096 — это 2.1 миллиарда параметров. При r = 8 обучаемых 8.4 миллиона, то есть 0.39%. Безубыточный ранг здесь 2048: до него ещё двести пятьдесят шесть удвоений ранга.

Обратите внимание на последний readout — вот где настоящая экономия. Обучаемый параметр стоит не одно место, а три: сам параметр, первый момент, второй момент. Для проекций внимания при r=8r = 8 это 9696 МБ против 2424 ГБ. Замораживание WW убирает три массива, а не один.

Почему малый ранг вообще работает

Здесь нужна честность: теоремы нет. Есть эмпирическое наблюдение, что поправка, нужная для адаптации к задаче, имеет быстро убывающий спектр, и есть теорема Эккарта–Янга из блока 1 (урок 090), которая говорит, что если спектр убывает быстро, то усечение по рангу — наилучшее из возможных приближений.

Посмотрите на спектр: если сингулярные значения падают быстро, ранг rr забирает почти всю норму; если равномерны — не забирает ничего.

поправка ΔW

ранга r · k = 2

ранг r 2

спектр сингулярных значений

ошибка по Фробениусу
0.378
хранимых чисел
66 / 256

Первые два-три сингулярных значения забирают почти всю норму. На такой поправке r = 8 не приближение, а практически точное представление — и ставка LoRA выигрывает.

Формулировка, которую стоит держать в голове: LoRA — это ставка на то, что поправка низкоранговая, а не факт про нейросети. Ставка обычно выигрывает, и известно, где она проигрывает: на задачах, требующих новых знаний, а не новой формы ответа, малого ранга не хватает, и это видно по тому, что качество упирается в потолок при росте rr вместо роста.

Три детали, без которых не работает

Инициализация B=0B = 0. Матрица AA заполняется случайно, а BB — нулями, поэтому в начале BA=0BA = 0 и модель в точности равна предобученной. Если бы обе были случайными, обучение начиналось бы с испорченной модели, и первые шаги уходили бы на восстановление.

Здесь стоит быть точным, потому что «одна из двух — нулями» звучит произвольно. Выпишем градиенты: при y=Wx+BAxy = Wx + BAx и δ=L/y\delta = \partial L/\partial y

LB=δ(Ax),LA=Bδx\frac{\partial L}{\partial B} = \delta (Ax)^\top, \qquad \frac{\partial L}{\partial A} = B^\top \delta x^\top

Отсюда видно, что нулём обязана быть ровно одна из матриц. Обе нулями — мёртвая точка: оба градиента обращаются в нуль и остаются нулями навсегда. Любая одна нулевая — работает: при B=0B = 0 на первом шаге двигается BB, а со второго и AA. Зеркальный вариант (A=0A = 0, BB случайно) тоже даёт ΔW=0\Delta W = 0 и тоже трогается с места, так что выбор между ними — не математика, а эмпирика: с AA случайным и B=0B = 0 на практике проходят большие learning rate.

Масштаб α/r\alpha/r. Без него величина поправки росла бы вместе с рангом, и learning rate пришлось бы подбирать заново под каждый rr. С ним rr и α\alpha разделены: ранг отвечает за выразительность, α\alpha — за силу.

Слияние. После обучения W=W+αrBAW' = W + \frac{\alpha}{r}BA считается один раз, и на инференсе никакой поправки нет вовсе — ни лишних слоёв, ни лишней задержки. Это главное отличие от адаптеров, которые добавляют в сеть новые блоки и потому платят на каждом прогоне.

Что LoRA не экономит

Тут распространено завышенное ожидание. Обратный проход обязан пройти через всю сеть: чтобы получить градиент по AA в первом слое, нужно протащить производную через все последующие слои, замороженные они или нет.

чтополное дообучениеLoRA
параметрывсе0.11%0.1{-}1\%
градиентывсетолько по AA, BB
состояние Adam2×2\times все2×2\times по AA, BB
активации для backwardвсевсе
проход назад через сетьцеликомцеликом

Экономия — в трёх верхних строках, и она огромна. В двух нижних — нулевая. Поэтому LoRA снимает ограничение по памяти под оптимизатор, но не делает шаг существенно быстрее: время по-прежнему уходит в прямой и обратный проходы.

Отсюда и QLoRA: если веса всё равно заморожены, их можно держать в 4 битах — они не обновляются, поэтому точность их представления нужна только для прямого прохода. Обучаемая часть при этом остаётся в fp16. Комбинация из двух независимых наблюдений, каждое из которых само по себе простое.

Итог

  • Замена одна: ΔW=BA\Delta W = BA ранга rr, и всё остальное — арифметика.
  • Безубыточный ранг nmn+m\frac{nm}{n+m} огромен по сравнению с используемыми, поэтому вопрос выбора rr — про качество, а не про память.
  • Экономятся градиенты и состояние оптимизатора; активации и время проходов — нет.
  • B=0B = 0 на старте, масштаб α/r\alpha/r, слияние после обучения — три детали, каждая из которых решает конкретную проблему.

Источники

Проверки

0 из 2
  1. Что экономит малый ранг

    Отметьте все верные утверждения о LoRA.

  2. Бюджет низкоранговой поправки

    Реализуйте lora_budget(in_dim, out_dim, count, rank, bytes_per_param) — верните [trainable, percent, break_even_rank, adam_megabytes]:

    • trainable = r(n+m)countr(n + m) \cdot \text{count} — параметры AA и BB для всех count матриц такой формы;
    • percent — их доля от nmcountnm \cdot \text{count}, в процентах;
    • break_even_rank = nmn+m\frac{nm}{n+m} — ранг, при котором поправка перестаёт что-либо экономить (не округляйте);
    • adam_megabytes — сколько мегабайт занимают обучаемые параметры вместе с двумя моментами Adam, то есть 3trainablebytes_per_param3 \cdot \texttt{trainable} \cdot \texttt{bytes\_per\_param}, делённое на 102421024^2.

    Проверить себя можно так: при rank, равном break_even_rank, percent обязан равняться ровно ста.

    функция lora_budget

    Загрузка редактора…

    Ctrl/⌘ + Enter