24 СЕНТЯБРЯ 2026 Г.
Ран Ран (Ran Ran), ведущий инженер-программист
OLMo 3, разработанная Институтом искусственного интеллекта Аллена (AI2), представляет собой передовую, полностью открытую языковую модель, обученную на основе современной архитектуры и многоэтапного рецепта обучения. Чтобы оценить возможности MaxText на Google Cloud TPU, наша команда поставила задачу воспроизвести OLMo 3 7B от AI2 с нуля. Мы выбрали OLMo 3, поскольку в ней сочетаются три свойства, которые редко встречаются вместе. Это мощная современная модель с 7 миллиардами параметров, обученная в условиях реального производственного масштаба. AI2 предоставляет практически полный цикл создания модели, включая данные, код, конфигурации, контрольные точки (чекпоинты), логи и результаты оценок. И наконец, у нас есть независимый эталон на PyTorch и GPU, с помощью которого мы можем протестировать MaxText и TPU.
Мы воспроизвели OLMo 3 7B от AI2 в MaxText на Google Cloud TPU (как предварительное обучение первого этапа, так и дообучение/отжиг второго этапа) и доказали совпадение по контрольным метрикам, а не только по кривой потерь (loss curve):
Основные моменты, каждый из которых подробно рассматривается далее в публикации:
- Конвертация модели из PyTorch в JAX. Архитектура OLMo 3 (блок с измененным порядком нормализации — reordered-norm, нормировка QK и соотношение скользящего и глобального внимания 3:1) была перенесена в MaxText и проверена с помощью тестов на логитное соответствие (logit-parity): конвертированный чекпоинт для шага 0 совпадает с эталоном HuggingFace с расхождением KL ≈ 1.5e-3 (что находится в пределах шума для «одной модели на разных фреймворках»), а при полном контексте в 8192 токена в формате bfloat16 обе реализации выдают одинаковый топ-1 токен в 98.75% случаев.
- Верификация, позволяющая выявлять реальные ошибки. Контрольные оценки помогли обнаружить ошибку в загрузчике данных, из-за которой MaxText казался лучше эталона; прирост на самом деле был результатом переобучения (memorization).
- Надежность при многонедельном запуске. Механизм сохранения и восстановления контрольных точек (checkpoint-and-resume) полностью воспроизводит процесс: контролируемое A/B-тестирование показывает Δ = 0.000 на каждом шаге после восстановления, а когда сбой хоста прервал запуск второго этапа на середине, восстановленный процесс дообучил 127 шагов с Δ = 0.000 по зарегистрированным потерям и перплексии.
- Изменение масштаба задачи обучения на лету. На шаге ~1.05M мы потеряли три четверти наших вычислительных мощностей; процесс был возобновлен на срезе (slice), составляющем четверть от прежнего размера, без изменения рецепта (тот же скрипт run_olmo3_7b_stage1.sh масштабирует размер батча на устройство для поддержания постоянного общего размера батча — GBS), при этом пропускная способность на одно устройство сохранилась в пределах 1% (≈100% сильного масштабирования, измерено в обоих направлениях).
- Смена поколения TPU в ходе выполнения рецепта. Для второго этапа тот же самый запускщик был направлен на v5p вместо Ironwood (изменился только тип устройства), что позволило стабильно удерживать MFU на уровне 57.4%.
- Эффективность оптимизации, которая полностью окупилась. 44.5% MFU на Ironwood для модели с 7B параметров благодаря коллективной выгрузке SparseCore, настройке пересчета градиентов (remat) и оптимальному разделению (sharding), что позволило сэкономить около трети вычислительного бюджета.
- Совместное проектирование для TPU: быстрее при том же качестве. Изменение формы внимания с 32 голов на 128-мерную голову на 16 голов по 256 (при идентичном количестве параметров и FLOP) работает на 12.4% быстрее, так как размер головы 256 полностью задействует матричные модули (MXU) 256×256 чипов Ironwood, а кривая потерь совпадает с оригинальной на отрезке до 120 млрд токенов (30 тыс. шагов). Это было побочное исследование; для основного воспроизведения была сохранена оригинальная архитектура.
Начиная с весов PyTorch шага 0 от AI2 и используя тот же базовый рецепт, запуск в MaxText повторяет опубликованную AI2 кривую потерь во всем диапазоне бюджета (~5.93 трлн токенов / 1.41 млн шагов) и успешно завершает первый этап. Мы даже упростили два элемента рецепта (использовали единое расписание косинусного изменения скорости обучения — cosine LR schedule — вместо двух, объединенных AI2, а также применили публично выпущенный набор данных; см. рецепт ниже), и соответствие результатов все равно сохранилось. В остальной части публикации рассказывается о том, как создавался, измерялся и — в одном поучительном случае — едва не был сфальсифицирован каждый из этих компонентов.
Зачем воспроизводить OLMo 3?
OLMo 3 — одна из немногих по-настоящему открытых языковых моделей пограничного класса: с открытыми весами, открытыми данными и полностью специфицированным рецептом обучения с публичным эталонным запуском на Weights & Biases. Успешное повторение этого независимо обученного прогона по контрольным метрикам, а не просто по кривой потерь, является весомым доказательством того, что стек MaxText (оптимизатор, функция потерь, конвейер данных, числовые вычисления) работает корректно, а не просто «выглядит так, будто обучается».
MaxText — это фреймворк для обучения языковых моделей на базе JAX/XLA, созданный специально для TPU. Перед нами стоял вопрос: можно ли достоверно воспроизвести рецепт PyTorch-on-GPU в среде JAX-on-TPU, добиться совпадения по важным метрикам, а не бит-в-бит, и как это доказать?
Рецепт OLMo 3 представляет собой трехэтапный учебный курс: общее предварительное обучение, промежуточное обучение (отжиг) и адаптация к длинному контексту. В этой статье рассматриваются этап 1 (предварительное обучение на ~5.9 трлн токенов) и этап 2 (промежуточное обучение). Оба этапа были полностью пройдены и сопоставлены с эталонами AI2. Этап 3 и постобучение (SFT/RL через Tunix) — это рецепты, которые мы подготовили, но еще не запускали.
Учебный курс предварительного обучения OLMo-3: этапы 1 (Ironwood) и 2 (отжиг на v5p), воспроизведенные в этой статье; этап 3 для длинного контекста (последовательность 65k, YaRN) и постобучение (SFT с последующим GRPO через Tunix) станут следующими. Учебный курс OLMo-3. Этапы 1 и 2 воспроизведены в этой публикации; на очереди этап 3 и постобучение.
Рецепт
OLMo 3 7B — это плотный трансформер (dense transformer) с 32 слоями и размерностью 4096, использующий несколько нестандартных решений: блок «переупорядоченной нормировки» (reordered norm), нормировку QK (QK-norm) и соотношение скользящего и глобального внимания 3:1. Конфигурация MaxText (olmo3-7b-pt.yml, используемая для этапов 1 и 2) полностью повторяет эти параметры:
Рецепт обучения отражает скрипт pretrain-1.py из OLMo-core; параметры, которые должны совпадать для выравнивания кривых:
Мы начали обучение с чекпоинта PyTorch для шага 0 от AI2, конвертированного в формат Orbax, поэтому MaxText стартует с тех же самых весов, что и эталон. Сама конвертация стала первой проверкой: прямой проход (forward pass) на конвертированных весах совпал с эталоном HuggingFace с расхождением KL ≈ 1.5e-3 и пересечением топ-10 токенов 9/10, что соответствует уровню шума «одной модели на разных фреймворках».
Зависит ли совпадение результатов от использования инициализации AI2? Судя по всему, нет. В качестве независимой проверки мы также провели обучение модели с нуля, используя собственную случайную инициализацию MaxText в течение примерно 50 тыс. шагов (3.5% от общего горизонта); ее потери при обучении близко следовали за опубликованной AI2 кривой, находясь чуть ниже нее. Это была точечная проверка потерь при обучении, а не полноценное воспроизведение, однако она предполагает, что совпадение не зависит от стартовых весов AI2.
Конвейер данных полностью повторяет OLMo-core: токенизация и конкатенация всех документов (с добавлением EOS между ними), нарезка на непересекающиеся фрагменты длиной 8192 токена, глобальное перемешивание индексов с фиксированным зерном (seed) и применение фильтра повторения n-грамм, который маскирует фрагменты с более чем 32 повторяющимися n-граммами. Этому соответствует тип датасета dataset_type=olmo_grain в MaxText (построенный на базе Grain).
Два намеренных отклонения от запуска AI2. (1) Расписание скорости обучения (LR schedule): изначально AI2 планировала обучение на ~5 трлн токенов и прямо в процессе расширила горизонт до ~5.93 трлн, поэтому график LR объединяет две косинусные кривые (это заметно в их публичном запуске на WandB); мы использовали единое косинусное расписание на весь горизонт. (2) Данные: мы обучаем модель на публично доступном наборе OLMo-3, в котором отсутствует небольшая доля токенов (<0.5% от общего бюджета, в основном это шарды s2pdf, отсутствующие в выпущенном списке файлов), которую видел внутренний запуск AI2. Оба отклонения являются нашими осознанными упрощениями, а не случайностями, и тем не менее MaxText совпадает с эталоном по всем контрольным показателям. Именно поэтому мы говорим о «воспроизведении в пределах межзапускового шума», а не о совпадении бит-в-бит (см. анализ KL в разделе §G).
Что нам пришлось построить
OLMo 3 не было в MaxText, когда мы начинали; в процессе воспроизведения всё нижеперечисленное было добавлено и передано в upstream. «Воспроизведите это самостоятельно» в конце этого поста — это конфиг, а не код.
- Сама модель: блок с измененным порядком норм (reordered-norm), QK-нормализация и схема внимания 3:1 (sliding/global) (#3004, #3112).
- Оптимизатор с пропуском шагов (skip-step), семантика которого полностью совпадает с olmo-core вплоть до скользящего стандартного отклонения с поправкой Бесселя (пропуск при 6σ в окне из 128 шагов) (#3490).
- z-loss (#3211) и маскирование затухания весов (weight-decay) для каждого параметра, позволяющее исключать эмбеддинги (#3280).
- Конвейер данных olmo_grain (#3749): чтения с произвольным доступом для предварительно токенизированных шардов, перемешивание глобального индекса с сидом и защитой в виде фингерпринта от незаметной подмены данных при перезапуске, фильтр n-граммных повторений и (начиная со второго этапа) чекпоинтинг состояния итератора Grain.
- Конвертация чекпоинтов HF↔Orbax с проверкой паритета логитов — инструмент, стоящий за каждым показателем «уровня шума фреймворка» в этом посте (#3112, #3832).
- Лаунчер для первого и второго этапов (скрипт запуска на основе переменных окружения + обертка XPK с функциями submit / monitor / resume_until_done) (#3886).
- Теги паритета TensorBoard (optim/step_skipped, perf/total_tokens), чтобы каждая метрика на дашборде W&B от AI2 имела аналог в MaxText для сравнения.
Сходится ли модель?
Главный результат представлен на едином графике: lm_loss первого этапа в MaxText по сравнению с опубликованной кривой AI2 в WandB, выровненной по шагам и усредненной с окном в 2 тыс. шагов. Примерно до 800 тыс. шагов обе кривые идут след в след с отклонением в пределах ±0.012; начиная примерно с 0.9 млн шагов, MaxText опускается ниже AI2 и больше не поднимается обратно — это первый признак ошибки в данных, подробно разобранной в следующем разделе.
Потери MaxText по сравнению с AI2 на первом этапе за 1,41 млн шагов, с разницей, сглаженной за 18 тыс. шагов, на нижней панели. Верх: кривые неотличимы в этом масштабе вплоть до хвостовой части. Низ: разница остается в пределах ±0.012 до примерно 800 тыс. шагов, уходит в отрицательную зону примерно с 0.9 млн по мере накопления повторов из-за ошибки в данных и резко падает после 1,25 млн, достигая −0.22 в сырых бинах по 2 тыс. шагов до применения сглаживания по 18 тыс. шагам, с которым отрисованы обе панели. Ничто из этого не влияет на потери на отложенной выборке (held-out loss) или нисходящую точность (downstream accuracy), как показывают следующие разделы. (Таблица по контрольным точкам: Приложение A; полные кривые сохранены как olmo_stage1_loss_curve.tsv.)
Но одна лишь кривая потерь — слабое доказательство: два прогона могут совпадать по потерям на обучении и расходиться во всем остальном, что действительно имеет значение. Поэтому мы проверили сходимость на четырех независимых поверхностях в шести контрольных точках на интервале в 915 тыс. шагов:
- Потери lm_loss на отложенной выборке C4: оценка только на прямом проходе (forward-only) на 16 млн токенов из отложенной выборки C4-en, идентичные батчи в обоих прогонах.
- Набор тестов lm-eval-harness из 8 задач: MMLU, HellaSwag, ARC-easy/challenge, OpenBookQA, PIQA, BoolQ, WinoGrande.
- Многодоменная перплексия на отложенной выборке: проход в стиле Paloma по веб-данным, новостям, энциклопедиям и смешанным доменам.
- KL-дивергенция на уровне токенов: расстояние между распределениями предсказания следующего токена на идентичных входных данных.
Как мы проводили измерения. Все оценки выполняются на парах чекпоинтов, выровненных по шагам: активный прогон MaxText в сравнении с чекпоинтом AI2 на том же шаге, т. е. его публичной ревизией HuggingFace allenai/Olmo-3-1025-7B@stage1-step{N}, сконвертированной в Orbax (конвертация воспроизводит эталон PyTorch с KL ≤ 1.8e-3, что соответствует уровню шума фреймворка). Для оценки используется стандартный lm-eval-harness (MMLU в 5 попыток, в остальных случаях — значения по умолчанию); σ — это стандартная ошибка (stderr) харнасса для каждой задачи, объединяемая в квадратуре для получения дельт.
К концу первого этапа все поверхности сходятся в том, что оба рецепта взаимозаменяемы:
Четвертая поверхность, KL-дивергенция на уровне токенов, — единственное значение, которое не является микроскопическим: в среднем 0.389 натов на идентичных входных данных, что примерно в 200 раз превышает уровень шума фреймворка. Этого и следует ожидать от двух независимых прогонов по одному и тому же рецепту (одинаковые совокупные навыки, разное распределение массы вероятностей), и именно поэтому мы говорим о «шуме от прогона к прогону», а не о совпадении бит-в-бит; детальный разбор приведен в приложении §G.
При этом разница в нисходящей точности (downstream accuracy) никогда не превышает ±0.005 макроусредненного значения по всем шести контрольным точкам, а знак меняется четыре раза — ровно то случайное блуждание, которого можно ожидать от двух честно выполненных прогонов, различающихся лишь генератором случайных чисел и числовыми погрешностями (таблица по контрольным точкам приведена в Приложении D):
Макроточность по 8 задачам (MaxText против AI2 в шести контрольных точках) с дельтой для каждой точки. Нисходящие возможности взаимозаменяемы в каждой контрольной точке. В отличие от потерь на обучении, точность никогда не расходится монотонно; дельта совершает случайное блуждание в пределах ±0.005 и заканчивается на отметке +0.0002.
Ошибка, которая выглядела как победа
Здесь начинается самое интересное. Начиная примерно с 0.9 млн шагов, потери MaxText на обучении опустились ниже потерь AI2 и больше не поднимались; после 1.25 млн они заметно ушли вниз — в среднем на −0.06, а на некоторых участках в несколько сотен шагов и вовсе на −0.25. Судя исключительно по графику потерь на обучении, можно было бы сделать вывод, что MaxText вырвался вперед.
Это было не так. Потери на отложенной выборке C4 для охватывающих чекпоинтов практически совпадали (Δ −0.004 на 1 000k, +0.003 в конце первого этапа), а точность на 1 350k слегка склонялась в пользу AI2. Потери на обучении падали, в то время как способность к генерализации не менялась. Это классический признак заучивания (memorization): модель видела некоторые последовательности более одного раза и показывала низкие потери на повторах.
Верх: потери MaxText на обучении падают до 1.63, в то время как у AI2 они остаются на прежнем уровне. Низ: дельта на отложенной выборке C4 остается близкой к нулю на протяжении всего процесса. Вся история на одном графике. Верх: в диапазоне 1.24M–1.41M потери MaxText на обучении многократно падают до 1.63 при усреднении по окну в 2 тыс. шагов (Δ −0.22; −0.25 на интервалах в сорок шагов), тогда как у AI2 они остаются плоскими. Это выглядит как оглушительная победа. Низ: дельта потерь на обучении (синий цвет) уходит в отрицательную область, но потери на отложенной выборке C4 (зеленые ромбы) никогда не покидают диапазон ±0.02. Потери на обучении упали, а генерализация — нет.
Причиной оказалась ошибка двойного шародования (double-sharding) в загрузчике данных Grain. Загрузчик OLMo для MaxText передавал ShardOptions(shard_index, shard_count) в загрузчик данных Grain DataLoader в то время, когда семплер индексов уже выполнял шародование внутри себя. Опция shard_options в Grain не просто записывает метаданные; она изменяет шаг (re-strides) потока индексов семплера. При shard_count=32 курсор данных продвигался в 32 раза быстрее, поэтому первый этап перестал быть одной чистой эпохой и превратился в процесc повторной выборки с возвращением Пуассона (Poisson(≈1) resample-with-replacement): примерно 37% корпуса так и не были увидени, 37% были увидены один раз, а 26% — два или более раз. Бюджет токенов остался прежним (~5.9 трлн реальных токенов), поэтому потери все еще глобально совпадали с AI2, но повторяющиеся экземпляры искусственно занижали потери на обучении именно там, где они возникали повторно.
Из этого можно извлечь два урока:
- Потери на обучении не являются доказательством сходимости. Единственная причина, по которой мы не опубликовали ложное заявление «MaxText превосходит эталон», заключалась в том, что мы обязались проводить оценку на отложенной выборке на каждой контрольной точке. Провал, вызванный заучиванием, невидим на отложенной выборке C4 и на всех 8 нисходящих задачах.
- Честное воспроизведение требует бюджета на ошибки. Исправление (grain.sharding.NoSharding(), позволяющее семплеру полностью контролировать шародование) сводится к одной строчке кода. На ее поиски ушли тестовый стенд A/B, юнит-тест, воспроизводящий расхождение при shard_count>1, и повторный запуск на железе для проверки.
При проверке исправления мы обнаружили второй, независимый баг: ошибку на единицу (off-by-one) при определении шага возобновления (resume-step). Номер директории чекпоинта равен N, но цикл обучения записывает директорию N после завершения итерации N, поэтому модель восстанавливалась до шага N+1, в то время как загрузчик данных возобновлял работу с батча N, повторно обучая один батч и затем постоянно выполняя работу на один шаг позади. С обоими исправлениями запуск с сохранением и возобновлением в точности воспроизводит непрерывный запуск: Δ = 0.000 в зарегистрированной функции потерь на всех 99 шагах. (В том тесте A/B использовалась загрузка данных с одним воркером; многоворкерный случай проявился на этапе 2, где мы его устранили; см. «Этап 2».) Для обоих багов написаны регрессионные тесты, которые падают на старом коде и проходят на исправленном.
Мы позволили текущему запуску этапа 1 завершиться как есть: он был выполнен на 85%, исправление не может раскодировать уже прочитанные данные, а перезапуск привел бы к потере примерно 1,2 млн шагов вычислений. Приведенная выше проверка показывает, что баг стоил нулевой наблюдаемой точности; исправление предназначено для будущих запусков.
Производительность и масштабируемость на Ironwood
Воспроизведение математических результатов — это лишь полдела; вторая половина — сделать это быстрым и поддерживать высокую скорость, когда кластер меняется у вас под ногами. За те недели, что занял запуск на 1,4 млн шагов, задание прерывалось, перепланировалось и изменяло размер не один раз, и стек должен был поглотить всё это, не меняя рецепт.
Выжимание максимальной полезности модели (MFU)
На Ironwood при модели 7B, размере батча на устройство 4, мы вышли на 44,5% MFU (510–513 Тфлопс/с/устройство) для стандартной архитектуры («вариант D», наше название из абляционного исследования в Приложении J). Изменение размерности головы, затрагивающее только форму тензора, дает более 49%; см. пункт «Размерность головы» ниже. Что изменило ситуацию в порядке влияния:
- Флаги Ironwood XLA + выгрузка на SparseCore: выгрузка коллективных операций (all-gather, 2D all-gather, reduce-scatter) на SparseCore, а также набор специфичных для v7x флагов XLA позволили поднять MFU с 41% до 44,5% без изменения потерь. (Полный список флагов см. в Приложении H.)
- Расширенная рематериализация: сохранение в чекпоинты проекций внимания и MLP (qkv_proj, q/k/v_proj, out_proj, mlpwi_0, mlpwo, context) уложилось в бюджет памяти активаций; добавление еще одной (mlpwi_1) привело к переполнению HBM на 18 ГБ.
- Внимание Splash + Tokamax с блоками по 2048 токенов.
- Ось шардинга не имеет значения при таком масштабе: чистый FSDP, 4-FSDP×32-DP и 8-FSDP×16-DP различались в пределах ~1,5 Тфлопс/с на 128 устройствах; чистый FSDP побеждает за счет простоты. Внутричиповый тензорный параллелизм (TP=2) оказался невыгоден: −1,6% MFU при половинном батче, OOM (переполнение памяти) при полном батче (FSDP=64 удваивает состояние весов на чип).
Масштабирование вверх и вниз, и почему это далось бесплатно
Самое полезное свойство стека JAX/XLA здесь заключается в том, что рецепт не зависит от топологии. Глобальный батч (512 экземпляров, 4,19 млн токенов/шаг) фиксирован; количество чипов, по которым он распределен — нет. (Техническое примечание: Ironwood содержит два устройства JAX на чип, поэтому срез 4×4×4 на 64 чипах предоставляет 128 устройств; мы упоминаем оба варианта.)
- Проверка масштабирования вверх: переход от среза со 128 устройствами к срезу с 512 устройствами (в 4 раза) при том же глобальном батче дал увеличение совокупной пропускной способности в 3,99 раза, то есть ≈100% сильного масштабирования в тесте на 1000 шагов. Сильное масштабирование — это сложное направление: каждое устройство теперь выполняет четверть работы за шаг, в то время как коллективные операции охватывают в 4 раза больше устройств, поэтому остается меньше доступных вычислений для перекрытия большего объема коммуникаций. Выгрузка на SparseCore в любом случае удерживала эти коммуникации вне критического пути.
- Масштабирование вниз в продакшене: на шаге ~1,05 млн мы потеряли три четверти нашей емкости, и запуск возобновился на срезе со 128 устройствами (вчетверо меньшем) при том же глобальном батче и без изменения рецепта. Пропускная способность на устройство сохранилась (~510–513 Тфлопс/с/устройство на обоих срезах); изменилось только реальное время выполнения шага (0,76 с → 3,05 с, ожидаемое увеличение в 4 раза).
Слева: совокупная пропускная способность масштабируется в 3,99 раза при переходе со 128 на 512 устройств. Справа: производительность в Тфлопс/с на устройство сохраняется при изменении размера.
Именно это позволяет долгосрочному запуску выживать в условиях конкурирующего за ресурсы кластера: берите любую свободную емкость, сохраняйте математику идентичной. На этапе 2 эта же идея была применена между поколениями TPU (см. ниже).
Автоматическое возобновление: выживание в многонедельном запуске
Запуск на 1,4 млн шагов обязательно будет прерван. Мы управляем им с помощью цикла resume_until_done, который автоматически отправляет задание заново при прерывании и возобновляет работу с последнего чекпоинта Orbax:
- Создание чекпоинта каждые 2000 шагов означает, что прерывание стоит максимум ~2000 шагов повторных вычислений (минуты на большом срезе). И мы добавили в цикл перезапуска нормальную экспоненциальную задержку: ранняя версия исчерпала MAX_RETRIES=50 из-за обратного давления в Kueue; настраиваемый параметр RETRY_BACKOFF_SECONDS (по умолчанию 300 с) позволил пережить многочасовые задержки планирования.
- TensorBoard на GCS является источником правды: логи kubectl видят историю только текущего пода. Каждая таблица в этом посте была сгенерирована на основе сохраненных в GCS событий TensorBoard, а не живых логов.
- Возобновление должно быть точным, иначе оно незаметно испортит запуск. Возобновление, которое читает не те данные или перезапускается на один шаг в сторону, выглядит нормально на кривой потерь, но это совсем не тот запуск, который вы думаете; это как раз те два бага загрузчика данных, о которых говорилось выше. После исправлений всё воспроизводится точно, и парный тест A/B ниже служит тому доказательством.
Багованное возобновление отклоняется от непрерывного запуска; исправленное возобновление дает ровно ноль на каждом шаге. Контролируемый тест A/B на 128 устройствах (64 чипа): возобновление из чекпоинта на шаге 100 по сравнению с непрерывным запуском. Ошибка на единицу (красный цвет) приводит к рассинхронизации данных с параметрами (среднее |Δ| 0.044, максимум 0.22) и никогда не сходится обратно. Исправление (зеленый цвет) дает ровно 0.000 на всех 99 шагах. Баг возобновления невидим на обычной кривой потерь; его можно увидеть только в таком парном сравнении.
Совместное проектирование аппаратного и программного обеспечения: бесплатное ускорение на 12%
Параллельно с воспроизведением мы провели абляционное исследование архитектуры, которое привело к серьезному успеху в области совместного проектирования аппаратного и программного обеспечения. OLMo-3 7B поставляется со стандартной конфигурацией из 32 голов запроса (query heads) и размерностью головы 128. Поскольку num_heads × head_dim = emb_dim = 4096, мы можем пожертвовать количеством голов в пользу ширины, изменив конфигурацию на 16 голов и размерность головы 256. Эта архитектурная корректировка сохраняет неизменными 7,298 млрд параметров и 1565 Тфлопс/шаг, но изменяет форму тензора так, чтобы она идеально гармонизировала с базовым «железом».
Изменение формы размерности головы: рост MFU с 44,2% до 49,6% (с 508 до 571 Тфлопс/с на устройство) при идентичных параметрах и FLOP.
Поскольку матричный умножитель (MXU) в Ironwood представляет собой систолический массив 256x256, стандартная размерность головы 128 оставляла половину массива простаивающей во время матричного умножения QK в механизме внимания. Изменение формы до размерности головы 256 идеально выравнивает размерность тензора до 256 в соответствии с аппаратным обеспечением, полностью устраняя циклы простоя вычислений. Это дает увеличение пропускной способности на +12,4% (571 против 508 Тфлопс/с/устройство, или 49,6% против 44,2% MFU), сохраняя при этом параметры и FLOP абсолютно идентичными. Это бесплатное ускорение, которое крайне рекомендуется внедрить перед запуском длительного обучения на стандартной конфигурации (нюансы кривой потерь и сидов обсуждаются в Приложении L).
Нюанс, заслуживающий отдельного абзаца
Тип данных оптимизатора тихий и дорогой. Установка weight_dtype=bfloat16 тихо понизила точность моментов m/v в алгоритме Адама через наследование mu_dtype, добавив +0.93 к потерям за 1000 шагов: мантисса bfloat16 длиной около 3 знаков отбрасывает часть каждого крошечного обновления на раннем этапе разогрева, и этот эффект накапливается. Сохранение weight_dtype=float32 (по умолчанию) сократило разрыв в 30 раз. Это был самый крупный момент из серии «почему это не совпадает» за весь проект.
Вычисления
Обучение этапа 1 потребовало около 77 тыс. чипо-часов Ironwood для вычислений на шагах (примерно 3200 чипо-дней), не считая контрольных точек, оценки и перезапусков. Чипо-часы — это неизменяемая единица: чипо-секунды на шаг не зависят от размера среза (0.76 с × 256 чипов ≈ 3.05 с × 64 чипа ≈ 195 чипо-с), в то время как реальное время зависит от конкретного среза. Основной объем вычислений выполнялся на срезе 4×8×8 (256 чипов / 512 устройств), где 77 тыс. чипо-часов эквивалентны примерно 12.5 дням; с учетом этапа после преемпшена на срезе из 64 чипов и времени ожидания в очереди общее календарное время составило несколько недель. При нашем бюджете в 30% MFU до запуска на те же токенизированные данные потребовалось бы примерно на 50% больше чипо-времени (~113 тыс. чипо-часов при том же расчете времени шага; в строке планирования 6·N·D из Приложения I указано около 100 тыс.), так что работа над производительностью сэкономила примерно треть. Этап 2 обошелся сравнительно недорого: около 5 тыс. чипо-часов v5p (~39 часов на срезе v5p-256 из 128 чипов; подробности в разделе «Этап 2» и Приложении I).
Этап 2: Середина обучения (отжиг)
Завершив этап 1, мы перешли к этапу 2: середине обучения, представляющей собой финальное затухание по расписанию warmup-stable-decay (WSD). Модель с этапа 1 проходит отжиг на смеси Dolmino 100B (высококачественные математические, кодовые, рассужденческие и отобранные веб-данные), в то время как скорость обучения линейно убывает от 2.0712e-4 до 0. Именно на этом этапе высококачественные данные OLMo-3 превращаются в прирост возможностей, поэтому полноценный стек должен также соответствовать этому этапу. Мы проверили рецепт по эталону середины обучения от AI2 (запуск zxv811e1, сгенерированный скриптом OLMo-core OLMo-3-1025-7B-midtrain.py); все гиперпараметры совпадают:
Инициализация Warm-Adam. AI2 выполняет мгновенный повторный разогрев с финального значения 3e-5 этапа 1 до 2.0712e-4 без этапа разгона и с параметром load_optim_state=True; второй момент Adam, загружаемый из чекпоинта, по всей видимости, сглаживает этот скачок. Мы воспроизводим это с помощью разовой модификации чекпоинта (сохраняем параметры + mu/nu, обнуляем шаг цикла и счетчик расписания скорости обучения), благодаря чему запуск возобновляет расписание на пике с уже разогретыми моментами, что соответствует конфигурации инициализации AI2. Потери на шаге 0 составили около 1.53, избежав всплеска «холодного старта».
Другое поколение TPU, тот же лаунчер: 57.4% MFU на v5p
На этапе 2 также произошло изменение поколений оборудования. Ресурсы Ironwood были задействованы в других задачах, поэтому мы направили идентичный скрипт запуска (включая флаги XLA, настроенные под Ironwood) на срез TPU v5p (v5p-256, 128 чипов), изменив лишь тип устройства XPK. Без какой-либо специальной настройки под v5p показатель MFU составил 57.4% (среднее значение 263 ТФлоп/с/чип при пиковом значении v5p в 459), что выше 44.5% на Ironwood для этапа 1, поскольку модель с 7 млрд параметров легче загружает более старый чип, чем чип с пятикратной пиковой производительностью. И этот показатель оставался стабильным: пропускная способность на чип находилась в диапазоне от 263.0 до 263.9 ТФлоп/с (от 25-го до 90-го процентиля всего прогона), что составляет разброс всего 0.4% за 47 684 шага. В сочетании с изменением размера на этапе 1 это полностью раскрывает тему портируемости: рецепт не зависит ни от топологии среза, ни от поколения TPU. Вы можете запускать обучение на любых свободных мощностях.
Происходит ли сходимость?
За все 47 684 шага разница в потерях при обучении по сравнению с данными AI2 составила в среднем +0.0044, причем она почти полностью обусловлена первыми примерно 8 тыс. шагов. Этот ранний разрыв не является следствием несоответствия рецепта: на начальном этапе две перетасовки обучались преимущественно на разных данных. Два случайных префикса из 8 тыс. шагов для смеси из 12.2 млн экземпляров пересекаются всего примерно на 17% (и всего на ~2% на 1 тыс. шагов, где разрыв достигает максимума). Как только области покрытия начинают пересекаться, разрыв сглаживается: каждое окно в 4 тыс. шагов после шага 12 тыс. находится в диапазоне от +0.0000 до +0.006, а на последней трети прогона среднее значение составляет +0.0007, что означает полное совпадение. Стоит отметить одну асимметрию измерений: наша функция перекрестной энтропии маскирует <|pad|> (следующий раздел), тогда как у AI2 она учитывается с почти нулевым весом. Это смещение искусственно занижает кривую AI2, поэтому значение +0.0044 в худшем случае является верхней границей для сопоставимого сравнения.
Ошибка возобновления только для этапа 2: сохраняйте в чекпоинт и итератор данных
Дальнейшие шаги
Оболочка
Скопировано
