Глава 13 из 14 30 мин
Своё дообучение
Как научить Ростка новой манере речи, не трогая его 17 миллионов чисел: две тонкие матрицы рядом с каждым слоем и несколько сотен примеров.
В этой главе
- понять, почему дообучение помещается в тонкую матрицу низкого ранга, и проверить это на картинках
- обучить LoRA-адаптер своими руками и влить его обратно в веса
- переключать Ростка между тремя голосами и выяснить, что каждый из них забыл
Росток уже умеет разговаривать. Отвечает на вопросы, рассказывает истории, вежливо объясняет, что понимает только по-английски. Но говорит он ровно так, как мы научили его в главе 12: просто, по-доброму и немного пресно. А если хочется, чтобы Росток разговаривал как пират? Или отвечал стихами? Или знал рецепты вашей бабушки, которых нет ни в одной книге?
Первое, что приходит в голову, — дообучить его ещё раз, тем же способом, каким мы учили его беседовать. Это сработает. Но посмотрите на цену: каждая новая манера речи — это ещё одна полная копия всех 17 миллионов чисел, а для обучения нужна память под каждое из них, причём в нескольких экземплярах. Для Ростка терпимо. Для модели на восемь миллиардов параметров это уже серверная.
Эта глава — о приёме, который сделал дообучение доступным всем, у кого есть ноутбук. Он называется LoRA, low-rank adaptation, «адаптация низкого ранга». Идея умещается в одну фразу: старые веса мы не трогаем вовсе, а рядом с каждой матрицей весов ставим тонкий обходной путь. Чтобы понять, почему этот путь может быть тонким, придётся немного поговорить об одной из самых красивых идей линейной алгебры — о ранге матрицы.
Дообучить всё целиком
Вспомним, что было в главе 12. Мы взяли базового Ростка и продолжили его учить, теперь на диалогах: та же функция потерь, тот же спуск, только данные новые. Меняться мог любой из 17 309 056 параметров. Это называется полным дообучением.
Во что оно обходится? Во время обучения оптимизатор хранит куда больше, чем сами веса. Каждому параметру нужен градиент, а Adam вдобавок держит для каждого два скользящих средних (моменты из главы 10). В обычной схеме со смешанной точностью набегает около 16 байт на параметр. Для Ростка это 17,31 млн × 16 байт ≈ 277 МБ — ноутбук и не заметит. Для Llama 3 8B та же арифметика даёт около 128 ГБ: больше, чем помещается даже в серверную видеокарту на 80 ГБ.
Есть и вторая цена. На выходе получается целая новая модель. Нужны пират и поэт? Ещё две копии. Сто клиентов, и каждому свой тон? Сто копий.
А ведь интуиция подсказывает, что новый стиль — изменение маленькое. Росток-пират по-прежнему знает грамматику, по-прежнему помнит, что Луна вращается вокруг Земли, по-прежнему ставит слова в правильном порядке. Меняется только манера: «arr», «matey», «be» вместо «is». Неужели ради маленькой перемены нужно 17 миллионов новых чисел? Чтобы ответить, сначала научимся измерять, насколько сложна перемена в матрице.
Сколько чисел нужно перемене
Каждый вес Ростка живёт в какой-то матрице — прямоугольной таблице чисел. Например, слой внимания в каждом блоке хранит матрицу 384 × 1 152, это 442 368 чисел. Таблицу чисел проще всего почувствовать как картинку: каждое число задаёт, насколько тёмен пиксель. Так что начнём с картинок.
Возьмём столбец чисел $u$ (по одному на строку) и строку чисел $v$ (по одному на столбец) и заполним таблицу по правилу «клетка = число её строки × число её столбца»:
$$W_{ij} = u_i \, v_j, \qquad W = u\,v^\top .$$Это внешнее произведение, и картинка из него получается совсем простая: каждая строка — один и тот же узор $v$, только умноженный на $u_i$, то есть темнее или светлее. Похоже на шотландку, где нити перекрещиваются под прямым углом. У такой матрицы ранг один, и чтобы её хранить, достаточно $m + n$ чисел вместо $m \cdot n$.
Если сложить несколько таких шотландок, можно нарисовать почти что угодно. Наименьшее число слоёв ранга один, которые в сумме дают матрицу в точности, называется её рангом $r$. Матрицу ранга $r$ можно хранить в $r(m + n)$ числах: $r$ столбцов и $r$ строк.
Поиграйте минуту — и бросятся в глаза три вещи. Росток узнаётся уже после горстки слоёв, хотя правая картинка хранит лишь малую долю чисел. Шотландка восстанавливается в точности на ранге 3: ровно так её и сделали, из трёх внешних произведений. А шум сжиматься отказывается: его сингулярные числа почти не убывают, каждый слой несёт примерно столько же, сколько предыдущий, и ошибка тает, только когда ранг подходит к полным 64. У случайных чисел нет структуры — значит, и сжимать нечего.
То же самое на numpy: SVD — одна строка, лучшая копия ранга $r$ — ещё одна. Запустите и сверьте последнюю строку вывода со строкой для ранга 8 — они должны совпасть, как и обещает теорема.
На ранге 10 ошибка падает до нуля: 1 600 чисел хранят всю картинку в точности. И это не случайность: каждая строка этой картинки — смесь всего десяти узоров. У круга, нарисованного на сетке, срезы бывают только восьми разных ширин — это восемь узоров солнца; волна добавляет ещё один, земля — ещё один. Ранг измеряет именно это: сколько независимых узоров на самом деле спрятано в таблице.
Ранг — это число независимых узоров в матрице. Матрица ранга r умещается в r·(m + n) чисел вместо m·n, а SVD находит лучшую такую копию для любой матрицы.
Теперь главный вопрос. Дообучение меняет веса: $\Delta W = W_\text{после} - W_\text{до}$. Эта перемена — тоже матрица. Какой у неё ранг? В 2020 году Армен Агаджанян с соавторами показали, что у дообучения на удивление низкая «внутренняя размерность»: языковую модель можно настроить на новую задачу, меняя на удивление мало чисел. Годом позже Эдвард Ху и его коллеги из Microsoft сделали практический вывод: раз перемена всё равно простая, её можно сразу искать в виде матрицы низкого ранга.
LoRA: обходной путь рядом с весами
Вот вся LoRA целиком. Слой Ростка вычисляет $y = Wx$. Мы замораживаем $W$: ни одно её число больше не изменится. А рядом ставим обходной путь из двух тонких матриц:
$$y = W x + \frac{\alpha}{r}\, B A\, x, \qquad A \in \mathbb{R}^{r \times d_\text{in}},\quad B \in \mathbb{R}^{d_\text{out} \times r}.$$$A$ сжимает вход до $r$ чисел, $B$ разворачивает их обратно до размера выхода. Их произведение $BA$ — матрица $d_\text{out} \times d_\text{in}$, той же формы, что и $W$, но ранг у неё не больше $r$. Учатся только $A$ и $B$. Для матрицы внимания 384 × 1 152 обходной путь ранга 8 хранит $8 \cdot (384 + 1152) = 12\,288$ чисел вместо 442 368 — это 2,8%.
Множитель $\alpha / r$ — ручка громкости. Благодаря ему смена ранга не меняет типичный размер поправки, и шаг обучения не приходится подбирать заново. В нашем lora.py $\alpha = 16$ и $r = 8$, так что обходной путь входит с множителем 2.
Есть одна тонкость: с чего начинать. LoRA заполняет $A$ маленькими случайными числами, а $B$ — ровно нулями. Тогда на первом шаге $BA = 0$, и модель в точности совпадает с чат-моделью: адаптер начинает с того, что ничего не делает, и учится уже отсюда. Но почему бы не обнулить обе? Посчитаем градиенты.
Всё это можно проверить и заодно посмотреть на процесс целиком — на игрушке. Роль замороженных весов $W$ здесь играет картинка Ростка. «Дообучение» — это костюм: цель — Росток, наряженный пиратом или поэтом. Перемена $\Delta W$ — сам костюм, и LoRA ранга $r$ должна найти её градиентным спуском, начав с $B = 0$.
Разные костюмы обходятся в разный ранг. На поэте берет и шарф — в основном горизонтальные и вертикальные полосы, и к рангу 8 костюм сидит почти идеально. Ремешок же пиратской повязки идёт по диагонали, а диагональ для низкого ранга — одна из самых дорогих вещей: в каждой строке своя тёмная точка на своём месте — вы видели это, когда рисовали сами. При том же ранге 8 ошибка у пирата всё ещё втрое с лишним больше, чем у поэта. И ещё одно: при любом ранге спуск приходит к пунктиру — к границе Эккарта — Янга. Две тонкие матрицы, обученные самым обычным спуском, сами находят лучшую перемену низкого ранга.
Что случилось бы, если бы и A, и B начинались с нулей?
В градиент по A входит B, а в градиент по B входит A. Если обе нули, каждый шаг спуска нулевой. Одна из двух должна начинаться со случайных чисел; LoRA выбирает A, чтобы произведение BA всё равно стартовало с нуля и модель в начале оставалась прежней.
Сколько это экономит
Теперь посчитаем на настоящих моделях. Наш lora.py ставит обходной путь рядом с каждой матрицей внутри блоков: общей матрицей Q K V во внимании, выходной проекцией внимания и всеми тремя матрицами сети прямого распространения (SwiGLU). Таблицу эмбеддингов он не трогает. Выберите модель, ранг и место для адаптеров:
Для Ростка LoRA — скорее удобство, чем необходимость: оба варианта помещаются в любой ноутбук. С ростом масштаба картина меняется. Полное дообучение Llama 3 8B требует больше сотни гигабайт, а с LoRA замороженные веса занимают 16 ГБ, и сверху добавляется лишь несколько сотен мегабайт на адаптер. Если же хранить замороженные веса в 4 битах (приём QLoRA), всё вместе ужимается до памяти обычного ноутбука. Именно так открытые модели дообучают дома.
В коде всё ещё короче, чем на словах. Весь слой LoRA из lora.py:
Три детали стоит сверить с математикой выше. A стартует случайной с масштабом $1/\sqrt{d_\text{in}}$, B — нулевой. forward сначала умножает на A.T и только потом на B.T: полная матрица $BA$ размером $d_\text{out} \times d_\text{in}$ так и не строится, и обходной путь стоит всего $r(d_\text{in} + d_\text{out})$ умножений на токен. А merge — одна строка, к ней мы сейчас вернёмся.
Обернуть модель так же просто: заморозить всё, а потом заменить каждый нужный слой его обёрнутой версией. Оптимизатор получает только список новых матриц.
Всё остальное в lora.py — цикл обучения из главы 12: те же render() и pack(), та же функция потерь, которая учитывает только слова самого Ростка.
Влить обратно
После обучения у нас есть замороженная $W$ и обходной путь $BA$. Можно было бы на каждом токене считать оба пути. Но обходной путь линейный, а значит, два пути склеиваются в одну матрицу:
$$W x + s\,BAx = (W + s\,BA)\,x = W' x .$$Считаем $W' = W + s\,BA$ один раз, записываем туда, где была $W$, а обходной путь выбрасываем. Остаётся обычный Росток: те же формы, те же 17 309 056 чисел, та же скорость. Поэтому Росток-пират в этой главе — обычный файл модели на 17,5 МБ, и браузерный движок загружает его, ничего не зная про LoRA.
lora.py сохраняет и то и другое: отдельно адаптер (adapter.pt, те самые 417 792 числа) и уже слитую модель. Сливать не обязательно. Если держать адаптеры отдельно, одна базовая модель в памяти может обслуживать сразу много стилей, подключая к каждому запросу свой маленький файл: так S-LoRA (Шэн и др., 2023) обслуживает тысячи адаптеров с одной видеокарты. Мы сливаем, потому что Росток маленький, а браузерный движок хочется держать простым.
Соберём теперь всю идею примерно в двадцать пять строк numpy, которые можно запустить: замороженный слой, спрятанная перемена ранга 2, которую адаптеру предстоит найти, обходной путь ранга 4, обученный ровно теми градиентами, что мы вывели, и слияние в конце.
Функция потерь падает с 89 до нуля, а после слияния слой выдаёт то же самое, что замороженный путь плюс обходной. Теперь поменяйте в строке d_in, d_out, r = 64, 48, 4 последнее число на 1 и запустите снова: функция потерь застрянет около 41. Искомая перемена имеет ранг 2, и обходной путь ранга 1 не может её выразить, сколько его ни учи.
LoRA не меняет модель, а учит рядом с ней тонкую поправку ранга r. После обучения поправка вливается в веса, и дообученная модель обходится ровно во столько же, во сколько исходная.
Три голоса Ростка
Пора за настоящее дело. Мы обучили два адаптера поверх чат-модели sprout-chat. Данные для каждого — 500 обычных диалогов Ростка, в которых каждый ответ переписан: для одного адаптера пиратской речью, для другого — в рифму. Вопросы остались прежними, переоделись только ответы. Настройки — те, что в lora.py по умолчанию: ранг 8, α = 16, AdamW с шагом обучения 0,003, девять проходов по данным, и функция потерь считается только на словах Ростка, как в главе 12. Число проходов мы подобрали по отложенным диалогам: их ошибка ниже всего примерно на седьмом-восьмом проходе, а потом снова растёт — 500 диалогов легко вызубрить. Это всего 34 шага для пирата и 41 для поэта: меньше минуты на Mac.
Задайте всем троим один вопрос. У них общее зерно случайности, так что ответы различаются только из-за весов.
Сколько пришлось выучить каждому адаптеру? Вот их настоящие кривые обучения. Толстая линия — ошибка на диалогах, на которых адаптер не учился (lora.py откладывает для этого каждый двадцатый); тонкая — ошибка на обучающих батчах.
Сравните, откуда кривые начинаются и насколько опускаются. Стартовая точка — это удивление чат-модели новой манерой речи, конечная — то, что от этого удивления осталось, когда адаптер ранга 8 усвоил несколько сотен примеров. Зазор между толстой и тонкой линиями — признак зубрёжки: если ошибка на обучающих батчах всё падает, а на незнакомых диалогах стоит на месте, адаптер заучивает конкретные фразы, а не стиль.
Что модель забывает
Чему бы нейросеть ни училась, учится она, меняя числа, которые уже чем-то были заняты. Когда дообучение тянет веса к новой задаче, старые навыки могут портиться. В крайнем случае это называют катастрофическим забыванием: сеть, обученная задаче Б, начисто забывает задачу А. Маклоски и Коэн описали его ещё в 1989 году на простейших сетях, и никуда оно с тех пор не делось.
Забывание можно измерить инструментом из главы 2 — удивлением. Дадим четырём версиям Ростка прочитать одни и те же тексты и посчитаем их среднее удивление, $-\ln p$ на токен. Модель, которая «разучилась» рассказывать обычные истории, удивится обычной истории сильнее.
На что смотреть? Если адаптеры выучили свои стили, в пиратской строке самым спокойным должен быть Росток-пират, а в стихотворной — Росток-поэт. Цена — в остальных строках: насколько сильнее каждый адаптер удивляется обычному ответу или обычной истории, чем чат-модель, из которой он вырос? Эта разница и есть забывание, измеренное в натах. На наших моделях оба адаптера свои стили выучили (в пиратской строке 2,93 против 4,14 у чат-модели, в стихотворной — 3,72 против 5,14), а забывание маленькое, но неравное: пират удивляется обычному ответу всего на 0,13 ната сильнее чат-модели, а обычной истории — не сильнее вовсе; поэт платит за обычный ответ 0,56 ната. Стихи уводят слова от обычной речи дальше, чем пиратские словечки, — сравните, где начинаются две кривые выше.
Ограничить его можно несколькими способами:
- Маленький ранг и мало шагов. Адаптер просто не в силах сдвинуть веса далеко.
- Немного старых данных в новой смеси. Несколько процентов обычных диалогов напоминают модели её привычную манеру; это называют повторением (replay).
- Вовремя остановиться. Следить за ошибкой на отложенных данных обоих видов и остановиться, когда старая начнёт расти.
- Беречь важные веса. Метод EWC (Киркпатрик и др., 2017) штрафует за изменение тех весов, которые больше всего значили для старой задачи.
И ещё одна честная оговорка. Дообучение, полное или LoRA, хорошо учит манере: тону, формату, длине, привычке рассуждать вслух. Добавлять знания оно умеет плохо, особенно в такую маленькую модель, как Росток, где все 17 миллионов чисел уже заняты английским языком и простыми историями. Если хотите, чтобы Росток «знал» бабушкины рецепты, реалистичный путь — показать ему рецепт прямо в вопросе.
Своё дообучение
Всё, что было выше, можно повторить на своих данных. Вот рецепт целиком.
1. Данные. Файл, где на каждой строке один диалог, в том же формате, что и в главе 12:
{"messages": [{"role": "user", "content": "how do birds fly?"}, {"role": "assistant", "content": "Arr, they flap their wings and push the air down, matey, and the air pushes them up, just like wind in a sail!"}]}
Для стиля хватает нескольких сотен диалогов: у наших адаптеров было по 500. Важнее объёма единообразие: если половина ответов пиратские, а половина нет, адаптер научится быть пиратом наполовину. И держите ответы на уровне Ростка: коротко и простыми словами. Разговаривать как профессор пятьсот примеров его не научат — он и слов-то таких почти не знает.
2. Обучение. Начинаем с чат-модели: стиль ложится поверх умения разговаривать.
# обучить адаптер и влить его в веса python lora.py --base runs/chat/ckpt_final.pt --chats my-style.jsonl --out runs/my-style # поговорить с результатом в терминале python chat.py runs/my-style/ckpt_final.pt # упаковать для браузера: int8, как в главе 11 python export.py runs/my-style/ckpt_final.pt sprout-my-style.bin
На Mac с чипом серии M lora.py сам возьмёт графический процессор через MPS, на машине с картой NVIDIA — через CUDA, а в Google Colab достаточно бесплатной среды с GPU (загрузите туда скрипты, data/tokenizer.json и контрольную точку чат-модели). Росток настолько маленький, что адаптер учится даже на обычном процессоре, просто медленнее. Сколько заняло наше обучение, показано в виджете с кривыми выше.
3. Проверка. Следите за проверочной ошибкой в log.jsonl: если она начинает расти, а ошибка на обучении всё падает, адаптер зубрит ваши примеры вместо того, чтобы учить стиль, — уменьшите --epochs. Поговорите с результатом на темы, которых нет в ваших данных: там и проявится забывание. Крутить стоит --rank (сравните 2, 8 и 32 по проверочной ошибке), --lr и --epochs.
После слияния Росток-пират пишет текст у вас в браузере. Насколько он быстрее или медленнее обычной чат-модели?
У W′ = W + s·BA та же форма, что у W. Движок выполняет те же умножения матриц, что и раньше, — различаются только числа внутри.
Росток сейчас
Росток научился менять манеру речи, не переучиваясь разговаривать заново. Каждый стиль — поправка ранга 8 из 417 792 чисел, обученная на 500 диалогах и влитая обратно в веса. Переключайте голоса прямо посреди разговора и смотрите, кто что помнит. В последней главе мы соберём все версии Ростка из этого курса в одном саду и посмотрим, куда двигаться дальше.