Может показаться, что размер батча стоит просто сделать как можно больше — до такой степени, пока будет влезать в доступную память, — и тогда каждый шаг обучения становится мощнее с точки зрения импакта на параметры.
Но давайте посмотрим на это чуть ближе. Для обучения можно определить две независимые оси, по которым нужно проводить оптимизации:
- Число шагов
- Тотальное количество обработанных примеров (батч × шаги)
Число шагов — это мера серийного времени. Каждый шаг идёт строго один за другим, и этот процесс не параллелится. Чем меньше таких шагов, тем быстрее в абсолюте мы можем обучить модель.
Тотальное количество обработанных примеров — это суммарный компьют, получается из размер батча × число шагов.
Каждый пример, прогнанный через модель, стоит каких-то вычислений (FLOPs). Меньше примеров в батче — меньше сожжённого компьюта.
И тут мы приходим к следующим трейд-оффам.
Если размер батча большой:
- то шагов нужно меньше, и мы экономим абсолютное время обучения;
- при этом за каждый шаг мы делаем вычисления для кучи примеров (как мы покажем ниже, лишних), и тратим компьют избыточно.
Если размер батча маленький, то:
- мы экономим компьют, не расползаясь на лишнее пространство за шаг;
- но шагов нужна тьма, и пока мы модель обучим, эксперимент проведём, вселенная погибнет.
И рыбку съесть, и на ёлку залезть не получается.
Тут важно понимать, что оси, которые мы рассматриваем, не приводимы. Можно попробовать сказать: и там, и там что-то тратится, это всё — какое-то время, а значит — какие-то вычисления и какой-то компьют, но это не так.
Давайте чуть подробнее:
- Компьют распараллеливается деньгами: чем больше видеокарт, тем больший батч может обрабатываться за один шаг, примеры внутри батча считаются независимо и параллельно на этих видеокартах.
- Шаги же обучения не распараллеливаются ничем: для шага N+1 нужен результат шага N — апдейтнутые веса, это всё происходит строго последовательно. Это время, которое не купишь бюджетом.
Компьют покупается, а шаги — нет. Даже с бесконечными картами каждый шаг ждёт предыдущий. Число обновлений весов до нужного качества является неустранимой последовательностью.
И тут напрашивается оффтопный вопрос: так ли это? Прямо вообще-вообще ничего сделать нельзя?
Можно, но всё это чего-то стоит и должно сравниваться отдельно, то есть «не при прочих равных».
- Можно сделать лучше оптимизатор, который до цели будет доходить за меньшее число шагов, то есть сходимость быстрее. Но опять: если вы берёте лучший оптимизатор, то в рассматриваемой нами задаче сравнения размера батча вы его берёте такой и там, и там. Поэтому на рассуждения не влияет.
- Можно применить асинхронные методы. Они не ждут друг друга. Разные воркеры считают градиенты и обновляют веса, не дожидаясь синхронизации. Но для этого они работают на слегка устаревших весах. И формально у нас получается разорванная последовательность. Из-за этого может появляться шум, а само обучение — сходиться хуже и нестабильнее. Мы покупаем этот параллелизм ценой качества шага. И на больших моделях это ведёт к большим проблемам. Так что тоже не катит.
- Можно попробовать сделать параллелизм по времени, но это область ресёрча, а не стабильного применения.
Так что абсолютное время обучения даже при большом размахе бюджета в ноль не схлопывается. И чем оно меньше, тем быстрее мы проводим каждый отдельный эксперимент, и тем быстрее у нас беклог проекта и сам проект едет.
Если у нас непомерно разросшийся батч экономит время, то он же и тратит компьют, требует больших затрат на железо; маленький батч экономит компьют, но тратит время.
Минимум компьюта стоит максимум времени, минимум времени стоит максимум компьюта.
[01]Как тогда правильно находить баланс размера батча?
Для этого рассмотрим величину шума градиента \(B_{\text{noise}}\). Она показывает отношение шума к сигналу в градиенте:
\[ B_{\text{noise}} = \frac{\text{разброс градиентов между примерами}}{\text{величина усреднённого градиента}} \]
Разные примеры тянут веса в чуть разные стороны.
Если этот разброс большой, а общее направление слабое, то усреднять много примеров (большой батч) реально полезно, потому что мы вытаскиваем сигнал из шума.
Если общее направление и так чёткое, то большой батч уже ничего не добавляет, и компьют на примеры тратится впустую.
[02]Что значит тратить компьют на примеры впустую
Когда пример попадает в батч, модель прогоняет его через forward и backward, и это стоит сколько-то FLOPs.
Впустую для нас означает, что эти вычисления почти не сдвинули веса, то есть обработка примера была сделана, а пользы от неё для весов почти не было.
Ключевой момент: в большом батче примеры усредняются в один общий градиент, и веса делают один шаг на весь батч. То есть 100 примеров в одном батче двигают модель ровно один раз, а не сто раз.
Сравним два способа потратить 8 миллионов примеров:
- Восемь шагов по 1M примеров: модель обновляется 8 раз, каждый раз учится на свежем усреднённом сигнале.
- Один шаг на 8M примеров: модель обновляется 1 раз. Все 8 миллионов слились в одно направление и дали одно-единственное движение весов.
Во втором случае мы прогнали в 8 раз больше примеров через модель (потратили в 8 раз больше вычислений), но обновление получили всего одно.
Работа сделана большая, а отдача маленькая, как от куда меньшего числа примеров.
Почему так происходит?
Усреднение по батчу гасит шум: батч из B примеров уменьшает шум примерно в B раз. Но это работает с убывающей отдачей. Пока батч маленький, а шум большой, добавить примеров может быть реально полезно, и каждый добавленный пример заметно уточняет направление.
Когда батч уже большой, шум и так задавлен, мы добавляем ещё примеры — они уточняют то, что и так уже точное.
Сигнал не становится существенно лучше, а вычислений ты потратил вдвое больше.
Поэтому на слишком большом батче ты тратишь непропорционально много компьюта на один шаг, а этот шаг почти не лучше, чем был бы на батче поменьше. Отдача на вложенные вычисления падает.
[03]Так как чему же равен оптимальный размер батча
Оптимальным будет являться размер, когда размер батча равен \(B_{\text{noise}}\). То есть размер батча сильно зависит от состояния весов модели в данный момент времени для данного датасета.
Покажем сначала, что значит оптимальный размер батча.
В этой точке оба трейд-оффа минимальны: они ровно вдвое больше минимально нужных. Если отходишь от этой точки, получаются взрывы по той оси, в которую отошли.
| Батч | Шагов (от минимума) | Данных (от минимума) |
|---|---|---|
| маленький (⅛ критического) | 9x | 1.1x |
| критический / оптимальный | 2x | 2x |
| большой (8× критического) | 1.1x | 9x |
Если мы берём батч поменьше, то экономим компьют (он почти на минимуме, 1.1x), но шагов теперь в 9 раз больше, время улетело вверх нелинейно.
Если мы берём батч побольше, то это быстро по шагам (время почти на минимуме), но компьюта тратится в 9 раз больше, чем было бы полезно.
[04]Почему именно \(B_{\text{noise}}\)?
Итак, \(B_{\text{noise}} = \frac{\text{разброс градиентов между примерами}}{\text{величина усреднённого градиента}}\) или \(B_{\text{noise}} = \frac{\operatorname{tr}(\Sigma)}{\lVert G \rVert^{2}}\).
Обозначим \(x = B / B_{\text{noise}}\).
Это отношение показывает, во сколько раз наш батч больше критического.
У нас есть закон насыщения. Прогресс за один шаг считается как
\[ \Delta L = \frac{\Delta L_{\max}}{1 + B_{\text{noise}}/B} \]
Тогда отсюда мы можем вывести такие соотношения:
\[ \frac{S}{S_{\min}} = 1 + \frac{1}{x} \]
Как отношение шагов ко времени. Мы смотрим, во сколько раз они больше минимума.
Здесь при огромном батче (x → ∞) шагов ровно минимум. При крошечном (x → 0) шагов бесконечно много.
И второе соотношение:
\[ \frac{E}{E_{\min}} = 1 + x \]
Соотношение примеров к компьюту: во сколько раз оно больше минимума.
При крошечном батче (x → 0) компьют ровно минимум. При огромном (x → ∞) компьют неограниченно возрастает.
Подставляем x = 1/8, x = 1, x = 8, и получаем ровно те 9× / 2× / 1.1× из таблицы выше.
Критический размер батча находится там, где взаимно минимальны обе оси, а не каждая из них по отдельности (потому что такой точки для обоих одновременно — не существует).
Другими словами, критический размер батча минимизирует произведение «время × компьют».
Перемножим:
\[ \frac{S}{S_{\min}} \cdot \frac{E}{E_{\min}} = \left(1 + \frac{1}{x}\right)(1 + x) = 2 + x + \frac{1}{x} \]
Получается, что выражение \(x + \frac{1}{x}\) минимально ровно при x = 1.
В любую сторону от x = 1 одно из слагаемых начинает расти быстрее, чем падает другое, и произведение раздувается.
Или мы это можем представить в виде уравнения гиперболы, если из каждой оси вычтем её минимум:
\[ \left(\frac{S}{S_{\min}} - 1\right)\left(\frac{E}{E_{\min}} - 1\right) = 1 \]
Видно, что перерасход времени и перерасход компьюта связаны так, что их произведение всегда есть единица, то есть они сами находятся в обратной пропорциональности друг от друга.
Устремляешь любой из них к нулю — другой улетает в бесконечность. Поэтому лучший компромисс как раз там, где они равны, то есть x = 1.
[05]Как это знание применять в жизни
- Критический размер батча — это не критический угол атаки, срыв потока после него не произойдёт. Если у вас есть ресурс, и его не жалко — берите батч побольше, если компьюта мало — поменьше.
Тут надо просто понимать все порядки законов, что с какой скоростью растёт, если мы выбираем отклоняться в ту или иную сторону, и просчитывать это, осознанно к этому подходить. - \(B_{\text{noise}}\) растёт по ходу обучения.
В конце общее направление становится слабым, шум начинает преобладать, и большой батч начинает оправдываться. Поэтому батч выгодно наращивать ближе к концу — так, например, GPT-3 доводил его примерно до 3.2M токенов. - \(B_{\text{noise}}\) почти не зависит от размера модели (при том же loss). Это удобный лайфхак: можно померить на маленькой модели, а затем перенести эту оценку на большую.
На сегодня всё.
В следующий раз будем говорить о том, как подбирать learning rate на более маленьких моделях, и затем его масштабировать для больших, а также Densing Law, который является аналогом закона Мура для LLM.