← Росток Своя LLM с нуля Словарик Код EN

Глава 10 из 14 35 мин

Обучение

Десять тысяч шагов, 330 миллионов токенов, три с половиной часа на одном Mac. Цикл обучения строчка за строчкой, оптимизатор, который дал нам фору, расписание скорости обучения — и замедленная съёмка того, как Росток учится писать.

В этой главе

  • прочитать train.py строчка за строчкой: случайные окна, ошибка, обратный проход, обрезка, шаг оптимизатора
  • понять, что AdamW и Muon делают с градиентом и почему Muon выиграл наш A/B
  • собрать расписание скорости обучения и посмотреть, как Росток шаг за шагом проходит путь от тарабарщины до историй

В конце прошлой главы Росток вырос в полный рост, но был совершенно пуст: 17,31 миллиона случайных чисел, и каждый токен удивлял его так, будто он выбирал наугад из всех 8 192. В этой главе мы наполним их смыслом. Всё обучение — один короткий цикл: показать модели кусок текста, измерить, насколько она удивилась, и чуть-чуть подвинуть каждое число туда, где удивление было бы меньше. Повторить 10 070 раз.

Прежде чем разбирать цикл, посмотрим, что он делает. По ходу обучения мы то и дело останавливались и просили Ростка продолжить четыре затравки. Вот что он писал.

Росток учится на глазах

Настоящие образцы из журнала основного обучения и ошибка на проверке в тот же момент. Двигайте ползунок, нажимайте на кривую или на ▶. Каждый образец — случайная выборка (температура 0,8, top-p 0,95), так что отдельный снимок может оказаться удачнее или неудачнее соседей.

Замедленная съёмка проходит через узнаваемые стадии. На нулевом шаге за затравкой идёт мешанина случайных токенов: у модели ещё нет никаких предпочтений. Через пару десятков шагов верх берут самые частые токены: точки, запятые, « the», « was», « and». Модель обнаружила, какие токены встречаются часто, — первое, что замечает любой, кто учит язык. Около 50-го шага появляются первые шаблоны — «Once upon a time, there was a girl named Lily», кавычки на своих местах, но грамматика ещё разваливается. Примерно к сотому шагу большинство предложений уже грамматичны, а диалоги записаны в правильном формате Имя: реплика, но история скачет с одного на другое. Через несколько сотен шагов истории держатся одной темы целый абзац. Дальше идёт медленная часть: меньше противоречий, всё длиннее связный текст, всё точнее выбор слов.

Следите и за числами над кривой. «Вариантов на токен» — это $e^{\text{loss}}$, перплексия из главы 2: из скольких равновероятных вариантов модель фактически выбирает. В начале их тысячи — даже чуть больше 8 192: случайная модель немного хуже честного угадывания наугад. К концу обучения — горстка. Это сжимающееся число и есть вся история этой главы.

Один шаг обучения

Вот сердце train.py — часть, которая выполняется 10 070 раз. Всё остальное в файле (аргументы, журнал, контрольные точки) — леса вокруг этих тринадцати строчек:

T = context_at(step, steps, args.context, args.seq_warmup) x, y = train.get(args.batch_tokens // T, T) with ctx: _, loss = train_model(x, y) loss.backward() norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) f = lr_factor(step, steps, args.warmup, args.cooldown) for opt in opts: for g in opt.param_groups: g['lr'] = g['base_lr'] * f opt.step() opt.zero_grad(set_to_none=True) seen += x.numel()

По строчкам:

  1. Длина окна. context_at сообщает длину окна на этом шаге: 128, 256 или 512 токенов. Зачем она меняется, расскажем ниже.
  2. Батч. train.get вырезает из train.bin 32768 // T случайных окон: 256 окон по 128 токенов или 64 окна по 512. К каждому окну прилагается его копия, сдвинутая на один токен, — ответы y. В батче всегда 32 768 токенов, так что каждый шаг задаёт модели сразу 32 768 вопросов «что дальше?».
  3. Прямой проход. Модель читает x и возвращает ошибку — среднее удивление по всем позициям всех окон. Строчка with ctx запускает её в bf16 — о нём чуть ниже.
  4. Обратный проход. loss.backward() — это обратное распространение из главы 4, применённое к 17 миллионам параметров: после него у каждого параметра есть .grad, наклон ошибки по нему.
  5. Обрезка. Если градиент подозрительно велик, его укорачивают (см. ниже).
  6. Шаг. Скорость обучения на этот шаг берётся из расписания, каждый оптимизатор двигает свои параметры, и градиенты обнуляются для следующего круга.

Обучение — это короткий цикл, повторённый десять тысяч раз: взять случайные окна, измерить удивление, подвинуть каждое число немного вниз по склону. Всё остальное — оптимизатор, расписание, bf16 — нужно, чтобы делать этот шаг быстрее и безопаснее.

Обрезка градиента

Иногда попадается необычный батч, градиент выходит огромным, и один неосторожный шаг может отбросить модель далеко от того места, где она была. Обрезка — ремень безопасности: если длина (норма) всего вектора градиента больше порога, вектор укорачивают до порога, сохраняя направление:

$$g \leftarrow g \cdot \min\left(1, \frac{1}{\lVert g \rVert}\right), \qquad \lVert g \rVert = \sqrt{\textstyle\sum_i g_i^2}.$$

clip_grad_norm_ заодно возвращает норму до обрезки, и мы её записываем. На кривых ниже видно, что в самом начале норма в несколько раз больше единицы: ремень действительно удерживает Ростка на первых хаотичных шагах. Примерно за первые полторы сотни шагов норма опускается ниже 1, и дальше обрезка почти не срабатывает.

bf16 и torch.compile: два бесплатных ускорения

bf16 («brain float 16») — 16-битный формат чисел с теми же 8 битами порядка, что у обычного 32-битного float, но всего с 7 битами мантиссы. Он покрывает тот же диапазон величин, от крошечных до огромных, только менее точно — две-три значащие десятичные цифры. Для обучения нейросетей такой точности обычно хватает, а умножения матриц в bf16 гоняют вдвое меньше байтов и работают быстрее. torch.autocast сам решает, какие операции можно перевести в bf16 (умножения матриц), а какие должны остаться во float32 (суммы, softmax, ошибка). Сами веса и состояние оптимизатора хранятся во float32.

torch.compile один раз смотрит на Python-код модели, строит из него граф операций и сливает мелкие операции в крупные программы для графического процессора (их называют ядрами), чтобы данные не гонялись между памятью и чипом ради каждого сложения. На нашем Mac это ускорило обучение примерно вдвое по сравнению с обычным («eager») режимом — без единого изменения в модели.

Как сделать шаг

После обратного прохода мы знаем градиент $g$ — направление, в котором ошибка растёт быстрее всего. Самый простой шаг — обычный градиентный спуск из главы 3:

$$\theta \leftarrow \theta - \eta\, g.$$

На практике у него две слабости. Во-первых, градиент одного батча шумный: он указывает примерно в нужную сторону, но от батча к батчу рыскает. Во-вторых, настоящий ландшафт ошибки похож на овраг: поперёк очень крутой, вдоль почти плоский. Скорость обучения, достаточно маленькая, чтобы не отскакивать от крутых стен, слишком мала, чтобы продвигаться по пологому дну.

Импульс превращает параметры в тяжёлый шар. Вместо текущего градиента мы шагаем по затухающей сумме прошлых: $m \leftarrow \beta m + g$, $\theta \leftarrow \theta - \eta m$ при $\beta = 0{,}9$. Зигзаги поперёк оврага гасят друг друга, а устойчивый толчок вдоль него накапливается.

Adam идёт дальше и даёт каждому параметру собственный размер шага. Кроме среднего градиента $m$ он следит за средним квадратом градиента $v$ и делит одно на другое:

$$m \leftarrow \beta_1 m + (1-\beta_1)\, g, \qquad v \leftarrow \beta_2 v + (1-\beta_2)\, g^2, \qquad \theta \leftarrow \theta - \eta\, \frac{\hat m}{\sqrt{\hat v} + \epsilon}.$$

Отношение $m / \sqrt{v}$ близко к $\pm 1$, когда градиент параметра устойчив, и близко к нулю, когда он всё время меняет знак. Поэтому параметр с устойчивым градиентом сдвигается примерно на $\eta$ за шаг, каким бы большим или маленьким ни был этот градиент. (Шляпки — поправка для первых шагов, пока $m$ и $v$ ещё близки к своим начальным нулям.) AdamW добавляет затухание весов — лёгкое притяжение каждого веса к нулю, которое применяется отдельно от градиента.

Устроим гонку четырёх оптимизаторов на трёх ландшафтах. Ландшафты двумерные, чтобы их было видно, но трудности настоящие: узкий наклонный овраг, изогнутая долина-«банан» и плато, где градиент почти нулевой.

Все четыре оптимизатора стартуют из одной точки, каждый со своей разумной скоростью обучения. Нижний график — высота каждого шарика в логарифмической шкале, пунктир — цель. «Muon» здесь — настоящий алгоритм из muon.py, применённый к матрице из одной строки: для такой матрицы он превращается в «шагни на фиксированную длину в направлении импульса».

SGD — обычный градиентный спуск — ползёт по оврагу и застревает на плато, где градиент крошечный, а шаги ещё крошечнее. Импульс разгоняется, но проскакивает цель. Adam и Muon не смотрят на величину градиента, только на его направление, поэтому пересекают плато на полной скорости. Эта независимость от масштаба — одна из главных причин, по которым современные сети учат адаптивными оптимизаторами. Но присмотритесь к Muon у цели: длина его шага постоянна, поэтому он первым добирается до дна, тут же перескакивает его и прыгает вокруг, пока спуск скорости в конце расписания не укоротит шаги. Adam тоже подпрыгивает, только меньше. Со спуском скорости мы ещё встретимся в разделе про расписание.

Muon: каждому направлению — одинаковый шаг

Adam смотрит на параметры как на длинный список независимых чисел. Но большая часть параметров Ростка сложена в матрицы, и у обновления матрицы есть своя структура. Здесь пригодится сингулярное разложение (SVD): любую матрицу можно записать как $G = U \Sigma V^\top$ — поворот, растяжение вдоль нескольких осей в $\sigma_1 \ge \sigma_2 \ge \dots$ раз и ещё один поворот. В градиенте весовой матрицы обычно несколько направлений имеют огромные $\sigma$, а остальные — крошечные. Обычный шаг тогда двигает веса в основном вдоль этих немногих главных направлений, а вдоль остальных — редких, но полезных — почти не двигает.

Muon (MomentUm Orthogonalized by Newton–Schulz — «импульс, ортогонализованный методом Ньютона — Шульца») сохраняет направления и выбрасывает растяжение: он заменяет сглаженный импульсом градиент ближайшей ортогональной матрицей,

$$G = U \Sigma V^\top \quad\longrightarrow\quad O = U V^\top,$$

у которой все сингулярные числа равны единице. Каждое направление обновления получает одинаковый шаг.

Считать SVD на каждом шаге было бы долго. Muon пользуется трюком — полиномиальной итерацией, которой нужны только умножения матриц. Начинаем с $X_0 = G / \lVert G \rVert$ (тогда все сингулярные числа не больше 1) и повторяем

$$X \leftarrow aX + b\,(XX^\top)X + c\,(XX^\top)^2 X.$$
Единичная окружность, преобразованная матрицей обновления (серым) и её ортогонализованной версией после $k$ итераций (цветом); рядом (на телефоне — ниже) — многочлен $f(\sigma)$ и путь обоих сингулярных чисел через него.

Вот оптимизатор, которым мы обучали Ростка, — весь muon.py, кроме шапки с импортами:

def orthogonalize(G, steps=5): """G: (..., rows, cols) -> the nearest semi-orthogonal matrices, batched.""" a, b, c = 3.4445, -4.7750, 2.0315 # tuned for fast convergence X = G.bfloat16() tall = X.size(-2) > X.size(-1) if tall: X = X.mT X = X / (X.norm(dim=(-2, -1), keepdim=True) + 1e-7) for _ in range(steps): A = X @ X.mT X = a * X + (b * A + c * A @ A) @ X if tall: X = X.mT return X.to(G.dtype) class Muon(torch.optim.Optimizer): def __init__(self, params, lr=0.02, momentum=0.95, weight_decay=0.0, split=None): # split: {param: n} treats a fused weight (like q, k, v) as n separate matrices self.split = split or {} super().__init__(params, dict(lr=lr, momentum=momentum, weight_decay=weight_decay)) @torch.no_grad() def step(self): for group in self.param_groups: # matrices of the same shape are orthogonalised together, in one batch by_shape = defaultdict(list) for p in group['params']: if p.grad is None: continue state = self.state[p] if 'momentum' not in state: state['momentum'] = torch.zeros_like(p) buf = state['momentum'] buf.lerp_(p.grad, 1 - group['momentum']) g = p.grad.lerp(buf, group['momentum']) # Nesterov look-ahead n = self.split.get(p, 1) by_shape[(p.size(0) // n, p.size(1))].append((p, g.view(n, -1, p.size(1)))) for (rows, cols), items in by_shape.items(): updates = orthogonalize(torch.cat([g for _, g in items])) scale = max(1.0, rows / cols) ** 0.5 i = 0 for p, g in items: u = updates[i:i + g.size(0)].reshape_as(p) i += g.size(0) if group['weight_decay']: p.mul_(1 - group['lr'] * group['weight_decay']) p.add_(u, alpha=-group['lr'] * scale)

Несколько деталей, важных на практике. Импульс использует «заглядывание вперёд» по Нестерову: обновление смешивает текущий градиент с буфером импульса. Матрицы одинаковой формы ортогонализуются вместе, одним пакетом, поэтому Muon обходится так дёшево. Объединённая матрица, которая считает запросы, ключи и значения за один раз, сначала делится на три части — по смыслу это три разные матрицы. В Muon идут только двумерные веса внутри блоков; таблица эмбеддингов и коэффициенты RMSNorm остаются у AdamW.

Итерацию легко проверить в numpy: возьмём «градиент», в котором одно направление больше чем в сто раз сильнее другого, и ортогонализуем его.

import numpy as np def orthogonalize(G, steps=5): a, b, c = 3.4445, -4.7750, 2.0315 # the same coefficients as muon.py X = G / (np.linalg.norm(G) + 1e-7) for _ in range(steps): A = X @ X.T X = a * X + (b * A + c * A @ A) @ X return X rng = np.random.default_rng(0) # a "gradient" of a 4x6 weight matrix where a few directions dominate G = rng.normal(size=(4, 6)) @ np.diag([10, 3, 1, 0.3, 0.1, 0.03]) U, S, Vt = np.linalg.svd(G, full_matrices=False) print('singular values of G: ', S.round(3)) print('after orthogonalize(G): ', np.linalg.svd(orthogonalize(G), compute_uv=False).round(3)) print('exact U @ Vt: ', np.linalg.svd(U @ Vt, compute_uv=False).round(3))

Наш A/B

Оптимизатор — не дело вкуса: его выбирают по измерениям. Перед основным обучением мы трижды обучили одного и того же Ростка на одних и тех же 10 миллионах токенов: с AdamW и с Muon при двух скоростях обучения.

Три коротких запуска полноразмерной модели, по 305 шагов. Всё одинаково, кроме оптимизатора для весовых матриц.

Результат однозначный: 2,393 у Muon против 2,709 у AdamW на проверке. Разница в 0,32 ната означает, что перплексия у Muon примерно в $e^{0{,}32} \approx 1{,}37$ раза ниже: его предсказания заметно менее «размазаны». Muon немного медленнее за шаг (29,7 тысячи токенов в секунду против 31,8 тысячи), но по качеству на токен выигрывает с большим отрывом. Удвоение его скорости обучения до 0,04 почти ничего не изменило — признак того, что он не капризен. Вопрос решён: основное обучение идёт на Muon со скоростью 0,02.

Что Muon делает с обновлением весовой матрицы, прежде чем применить его?

Это и есть ортогонализация: $U\Sigma V^\top \to UV^\top$. Поэлементное деление — это Adam; Muon смотрит на матрицу целиком и выравнивает её направления, чтобы редкие направления не тонули в главных.

Расписание скорости обучения

Скорость обучения $\eta$ не остаётся постоянной. В нашем расписании три части:

  • Разогрев, 200 шагов: скорость линейно растёт от почти нуля до пика. В начале веса случайны, градиенты большие и хаотичные, а средние оптимизаторов ($m$, $v$ и буфер Muon) ещё пусты. Полноразмерный шаг в этот момент может забросить модель туда, откуда она будет выбираться тысячи шагов.
  • Плато: большую часть обучения скорость держится на пике. Большие шаги быстро двигают модель, но шум батчей не даёт ей улечься на дно долины — она всё время дрожит вокруг него.
  • Спуск, последние 30%: скорость линейно падает до нуля. Шаги мельчают, шум усредняется, и модель оседает на дно. Ошибка на этом участке обычно заметно снижается.

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

Сверху — скорость обучения на 2 000 шагах игрушечной задачи. Снизу — её ошибка в логарифмической шкале (скользящее среднее по 10 шагам); серым — прошлый запуск для сравнения. Шум один и тот же во всех запусках, пока вы не нажмёте «Другой шум».

Несколько опытов, которые стоит поставить. Уберите разогрев — первые же шаги сбросят шарик с обрыва. Уберите спуск («Постоянная») — ошибка так и останется на уровне шума. Поднимите пик — сначала прогресс быстрее, но с какого-то момента побеждает шум, а ещё чуть дальше не спасает даже разогрев. Сравните линейный спуск с косинусом: итоги близки, и это совпадает с опытом на настоящих моделях.

Вот две функции из train.py, которые задают расписание (без изменений), и то, что они дают для обучения Ростка:

def lr_factor(step, total, warmup, cooldown): """Warm up, hold, then decay linearly to zero over the last `cooldown` share.""" if step < warmup: return (step + 1) / warmup decay_start = total * (1 - cooldown) if step < decay_start: return 1.0 return max(0.0, (total - step) / (total - decay_start)) def context_at(step, total, context, warm): """Sequence-length warm-up: short windows first, full length later.""" if not warm: return context if step < 0.1 * total: return max(64, context // 4) if step < 0.3 * total: return context // 2 return context total = int(330e6 // 32768) # 10,070 steps for step in [0, 99, 199, 1006, 1007, 3021, 7048, 8500, 10069]: T = context_at(step, total, 512, True) print(f'step {step:5d} lr x {lr_factor(step, total, 200, 0.3):.3f} window {T} windows per batch {32768 // T}')

Вторая функция — разогрев длины окна. Первые 10% шагов Росток читает окна по 128 токенов, до 30% — по 256 и только потом полные 512. Число токенов в батче остаётся тем же, просто окон больше и они короче. Короткие окна дешевле — стоимость внимания растёт как квадрат длины, — а в начале модель всё равно занята локальными закономерностями: словами, пунктуацией, грамматикой внутри предложения. Дальние связи она учит позже, когда до них дорастёт.

Настоящее обучение

Теперь основной запуск: 10 070 шагов, 330 миллионов токенов, Muon и AdamW, расписание выше, всё на одном Mac с M4 Pro. Он занял 3 часа 41 минуту (13 289 секунд), в среднем 24,8 тысячи токенов в секунду вместе с проверками, а итоговая ошибка на проверке — 1,554 ната на токен. Каждые 20 шагов скрипт записывал ошибку на батче, скорость обучения, длину окна и норму градиента; каждые 250 шагов — измерял ошибку на проверочном тексте (каждый раз на одних и тех же 20 батчах окон по 512 токенов).

Серым — ошибка на каждом обучающем батче, терракотовым — на проверочном тексте. Фон показывает три длины окна, пунктир — начало спуска скорости. Ведите по графику пальцем или мышью, чтобы прочитать значения; нижняя панель переключается между скоростью обучения, длиной окна и нормой градиента.

На что смотреть:

  • Первая сотня шагов уводит ошибку примерно с 9 ниже 4: модель открывает частоты токенов и простейшие пары. Переключитесь на логарифмическую шкалу: на ней всё обучение видно лучше, а длинная середина превращается в пологий, почти прямой склон.
  • Кривые обучения и проверки идут вместе. Модель почти никогда не видит один и тот же текст дважды (вспомните мозаику из прошлой главы: четыре пятых корпуса остаются непрочитанными), поэтому ей нечего зазубривать. В таком режиме нет переобучения, и ошибка на проверке — честная мера прогресса.
  • Смена окон оставляет след. Проверку всегда проводят на окнах по 512 токенов, а первые 10% обучения Росток не читал ничего длиннее 128. Экзамен был труднее уроков, и сразу после перехода на длинные окна ошибка на проверке падает заметно быстрее. Ошибка на обучении при каждой смене тоже скачком снижается: с более длинным контекстом следующий токен угадывать легче.
  • Норма градиента (нижняя панель) начинается намного выше порога обрезки и опускается ниже него примерно за первые полторы сотни шагов.

В логарифмической шкале середина кривой — длинный пологий склон: каждое десятикратное увеличение числа шагов снимает похожий кусок ошибки, с каждым разом чуть меньший. Такие кривые описывают степенным законом: $\mathcal{L}(D) \approx E + A \cdot D^{-\alpha}$, где $E$ — часть ошибки, которую не убрать никаким количеством данных. Каждое десятикратное увеличение данных снимает одну и ту же долю того, что осталось сверх $E$ (на графике с логарифмами по обеим осям $\mathcal{L} - E$ — прямая). Это закон масштабирования, который Kaplan et al. (2020) нашли для моделей самых разных размеров и который уточнила Chinchilla в 2022 году. Благодаря ему ошибку огромной модели можно предсказать по серии маленьких, прежде чем тратить миллионы на её обучение. Мы вернёмся к нему в главе 14.

Почему ошибка Ростка на проверке так точно следует за ошибкой на обучении, без признаков переобучения?

Переобучение — это зазубривание конкретных примеров, а для этого их надо видеть многократно. Росток не успевает пройти свои данные даже один раз, так что каждый батч для него — почти всегда новый текст, такой же незнакомый, как проверочный. Модели, которые много эпох учатся на маленьком наборе, — другая история.

Контрольные точки бок о бок

Во время обучения мы несколько раз сохраняли веса Ростка. Дадим всем версиям одну и ту же затравку и одно и то же случайное зерно и сравним. Каждая контрольная точка — полноценная модель на 17,5 МБ, поэтому при первом запуске они скачиваются одна за другой.

У всех контрольных точек одно зерно, температура 0,8 и top-p 0,95, так что различия идут только от весов. Ниже — все модели курса на одном и том же отложенном тексте, в натах на букву.

Лестница внизу показывает место Ростка среди всех моделей курса. Буквенная биграмма из главы 1, буквенные сети, сеть на токенах, трансформеры в один и четыре блока — каждая следующая модель немного снижала ошибку. Росток на шаге 200, всего после 6,5 миллиона токенов, уже обгоняет сеть на токенах. К шагу 1 000 он почти догоняет четырёхблочный трансформер из главы 8, а остаток обучения уводит его гораздо дальше: сравните ошибку на проверке на шаге 1 000 и в конце на кривых выше. Идеи всё время были те же; понадобились размер, данные и хорошо собранный цикл обучения.

Росток сейчас

Это результат предобучения — базовая модель. Она прочитала треть миллиарда токенов и продолжает любой текст в духе своего корпуса. Пишет детские истории с героями, у которых есть имена, держит диалог в формате Имя: реплика, знает, как история начинается и заканчивается. Но на вопросы не отвечает: спросите её о чём-нибудь — и она просто продолжит писать, будто ваш вопрос — строчка из рассказа или пьесы: новые реплики, слова рассказчика, выдуманные собеседники, но не ответ вам. А слова она выбирает, бросая кости, которые мы для неё настроили. В следующей главе мы разберём эти кости — температуру, top-k, top-p — и увидим, почему жадный выбор застревает в петлях.

Главы

  1. 0 Знакомство
  2. 1 Считаем буквы
  3. 2 Мера удивления
  4. 3 Градиентный спуск
  5. 4 Обратное распространение
  6. 5 Эмбеддинги
  7. 6 Токены
  8. 7 Внимание
  9. 8 Трансформер
  10. 9 Корпус
  11. 10 Обучение
    1. Росток учится на глазах
    2. Один шаг обучения
    3. Как сделать шаг
    4. Расписание скорости обучения
    5. Настоящее обучение
    6. Контрольные точки бок о&nbsp;бок
    7. Росток сейчас
  12. 11 Выборка
  13. 12 Разговор
  14. 13 LoRA
  15. 14 Что дальше