Порт модели на AI-ускоритель прошёл все тесты — и всё равно выдавал мусор

Порт модели на AI-ускоритель прошёл все тесты — и всё равно выдавал мусор

Модель скомпилирована. Тесты показывают SEQ_MATCH: True — вывод на устройстве побайтово совпадает с эталоном на CPU. Можно деплоить. Вот только эталон тоже сломан, и оба генерируют бессмыслицу. Эта история — о пяти ловушках при портировании 30-миллиардной языковой модели на специализированный AI-ускоритель, где каждый пройденный тест оказывался иллюзией корректности.

Контекст: Gemma-4 на Inferentia2

Разработчик под ником xbill портировал все пять вариантов Google Gemma-4 на AWS Inferentia2 — специализированные чипы Amazon для инференса нейросетей. Модели: E2B, E4B, 12B (компактные), 26B-A4B (MoE, смесь экспертов) и 31B (плотная, dense). Оборудование — инстанс inf2.24xlarge с 6 чипами Neuron, 12 ядрами и 384 ГБ памяти. Фреймворк — torch-neuronx 2.8.0 и neuronx-distributed.

Первые три модели (до 12B параметров) портировались относительно гладко: архитектурные особенности вроде Per-Layer Embeddings и cross-layer KV-sharing усложняли работу, но масштаб был управляемый. 31-миллиардная модель отбросила все эти сложности и подкинула проблему другого рода — масштаб. А заодно и ловушку, которая стоила дней отладки.

Ловушка первая: рецепта, работавшего на 12B, не существует для 31B

Для моделей до 12B автор вручную управлял тензорным параллелизмом через parallel_model_trace из neuronx_distributed — сам контролировал ранги, KV-кэш, весь цикл компиляции. На 31B этот подход умер двумя способами одновременно.

Первая попыка: все 8 рангов трейсятся в одном процессе. Модель весит около 15 ГБ на ранг в bf16, и восемь одновременных графов компиляции в fp32 съели 300+ ГБ памяти и убили машину с 384 ГБ. Это не модель большая — это трейсер восьмикратно параллелен.

Вторая попыка: последовательная компиляция рангов через max_parallel_compilations=1. Память упала до 83 ГБ, один ранг дошёл до статуса Compiler PASS — и multi-rank rendezvous завис. Воркер заблокирован на pipe_read, процессор простаивает, 25 минут тишины.

Мораль: при 30 миллиардах параметров ручное управление трейсером само становится багом. Решение — NxD ModelBuilder, который компилирует один ранг и загружает веса через checkpoint_loader. Компиляция одного ранга, шардирование весов на все 8, пиковая нагрузка 182 ГБ из 384 доступных. Время сборки — 39 минут, результат — 108 ГБ скомпилированных neff-файлов.

Ловушка вторая: два типа внимания, один из которых невозможно шардировать

Gemma-4 31B чередует два типа attention-слоёв. Скользящее внимание (sliding): 50 слоёв, 32 Q-головы, 16 KV-голов, head_dim 256. Глобальное внимание: 10 слоёв, 32 Q-головы, 4 KV-головы, head_dim 512. При тензорном параллелизме TP=8 нельзя разшардировать 4 KV-головы на 8 рангов — это физически невозможно.

Правило, которое работает: если количество KV-голов делится на TP без остатка — шардировать всё (q_proj, k_proj, v_proj — ColumnParallelLinear, o_proj — RowParallelLinear). Если не делится — оставить слой целиком на каждом ранге (replicate). Поскольку шардированные слои используют row-parallel для o_proj (all-reduce обратно в полное скрытое состояние), а реплицированные слои принимают и возвращают полное состояние, две компоновки свободно компонируются через residual stream.

Сигнал к действию прост: nkv < TP означает «реплицируй, не шардируй». Это правило уместилось в четыре строки кода, но его отсутствие стоило нескольких дней.

Ловушка третья: буфер — не параметр, и компилятор об этом не предупредит

Gemma-4 масштабирует выход каждого слоя через обучаемый layer_scalar. В коде модели он зарегистрирован как register_buffer, а не как параметр. ModelBuilder при загрузке шардированного чекпоинта перемещает только параметры. В итоге все 60 layer_scalar тихо дефолтнулись в 1.0, и шестьдесят слоёв ошибки, умноженной на константу, сложились в шум.

Косинусное сходство с референсом? Примерно 0. Никаких сообщений об ошибках. Модель просто работала неправильно, и единственный индикатор — качество вывода, которое выглядит как случайный шум.

Решение: читать буферы напрямую из safetensors и копировать их вручную после загрузки. Строка lyr.layer_scalar.copy_(lsv[i]) в цикле по всем слоям — и модель оживает. Мораль, прибитая к монитору автора: всё, что зарегистрировано через register_buffer — послойные скаляры, некоторые нормализации — требует явного копирования. Иначе модель тихо использует дефолты из конфига.

Ловушка четвёртая: идеальный SEQ_MATCH — и оба вывода мусор

После компиляции — финальная проверка. Прогоняется промпт на CPU в fp32 и на устройстве. Результат:

CPU: <start_of_turn>model\n<start_of_turn>model\n<start_of_turn>model...
Device: <start_of_turn>model\n<start_of_turn>model\n<start_of_turn>model...
SEQ_MATCH: True

Устройство воспроизвело CPU-референс токен-в-токен. Компиляция численно безупречна. Вот только вывод — мусор. Модель зацикливается на маркерах turnoв.

Ключевая подсказка: когда устройство совпадает со сломанным референсом, баг не в ускорителе. Он выше обоих — в том, что они разделяют. В данном случае — в промпте.

Ловушка пятая: токенизатор тоже врёт

Два факта о снапшоте Gemma-4, которые сложились в катастрофу. Во-первых, чат-шаблон поставляется отдельным файлом chat_template.jinja (18,7 КБ, с scaffold для thinking-токенов) — он не встроен в tokenizer_config.json. Поэтому apply_chat_template() выбрасывает ошибку «no chat template set».

Во-вторых, маркеры хода в Gemma-4 — это <|turn> (token id 105) и <turn|> (token id 106), а не <start_of_turn>. «Очевидный» фоллбэк — написать «<start_of_turn>user\n...» вручную и токенизировать — делает вот что:

"<start_of_turn>" → ['<', 'start', '_', 'of', '_', 'turn', '>'] — семь букв вместо одного токена. А convert_tokens_to_ids("<start_of_turn>") возвращает 3 — unknown token.

Модель получает сломанный промпт, предсказывает < как наиболее вероятный следующий токен, и зацикливается. CPU-референс, скормленный тем же сломанным промптом, делает ровно то же самое. SEQ_MATCH = True — всё время, на полной бессмыслице.

Исправление: загрузить и применить настоящий шаблон из chat_template.jinja. После этого — «The capital of France is Paris.» на обоих устройствах, prefill ~115 мс.

Параллельная история: 128 экспертов, которые голосовали не за тех

Пока 31B боролась с масштабом, модель 26B-A4B (Mixture of Experts) подкинула проблему другого измерения. Название «A4B» означает ~4 миллиарда активных параметров на токен — звучит как компактная модель. Это ложь, которую не простит ваш бюджет памяти. Все 128 экспертов (~49 ГБ) обязаны находиться в HBM одновременно. Top-8 маршрутизация сокращает вычисления, но не footprint. Поэтому бюджетировать нужно под 26B, а не под 4B — и это требует того же 192 ГБ, что и вдвое большая плотная 31B.

Архитектура оказалась страннее, чем «MLP заменён на MoE». Каждый из 30 слоёв запускает плотный MLP параллельно с 128-экспертным MoE, потом комбинация проходит через четыре feed-forward layernorm. Router: собственный RMSNorm + scale + per_expert_scale, затем softmax(128, fp32) → top-8 → renormalize → умножить на per_expert_scale. Эксперты: fused gate_up_proj [128,1408,2816] + down_proj [128,2816,704], где forward — это sparse gather/scatter loop через torch.where и index_add_.

Проблема: этот sparse loop зависит от данных и не трейсится в статический Neuron-граф. Решение — хак, который одновременно элегантен и расточителен. Вместо того чтобы собирать только top-8 экспертов на каждый токен, вычисляются все 128 экспертов, каждый умножается на свой router weight (который равен нулю для невыбранных), и всё суммируется. Поскольку невыбранный эксперт даёт 0 × expert(x) = 0, это математически идентично HF sparse top-8, но представляет собой фиксированную последовательность matmuls, которую компилятор обожает. Расточительно по FLOP (считаются 128 экспертов, используются 8), корректно в каждом бите, и трассируемо.

Всё это сворачивается в два стандартных параллельных линейных слоя, которые ModelBuilder уже умеет шардировать: gate_up_proj становится одним ColumnParallelLinear (ранг r получает экспертов 16*r…16*r+15), down_proj — одним RowParallelLinear (вход зашардирован → all-reduce). Никакого кастомного 3D-шардирования параметров.

Математика проверена на CPU: MAXDIFF ≈ 2e-6, косинус 1.0. Трейсинг на устройстве: успешен, первый в истории MoE-трейс на Neuron — 30 MoE-слоёв. Запуск: вывод пустой. Первый токен = end-of-turn, немедленно. CPU-референс: «The capital of France is Paris.» ✅ Математика: ✅ Шардирование: ✅ Устройство: ❌

Когда математика доказана, шардирование доказано, а устройство всё равно неправильно — баг в разрыве между «как я думаю, работают примитивы NxD при трейсинге» и «как они работают на самом деле». Конкретно: для взвешивания экспертов ранга r по router я рассеивал плотную матрицу весов через scatter_to_tensor_model_parallel_region. Эта функция выбирает срез через get_tensor_model_parallel_rank() — Python int. Но ModelBuilder компилирует один ранг и реплицирует граф на все 8 при рантайме. Значит, трейс встроил срез ранга 0 (Wd[:, 0:16]) в граф, и при рантайме каждый ранг взвешивал экспертов 0–15, пока его gate_up вычислял своих экспертов (16–31, 32–47, …). Полное рассогласование → мусор → модель уверенно завершает ход.

Auto-sharded ColumnParallel/RowParallel не имеют этой проблемы, потому что их веса предварительно нарезаны при загрузке — граф forward ранг-агностик. Standalone collective на рантайм-активации — имеет. Исправление: зарегистрировать SPMDRank-модуль ([1] int32 параметр, загружаемый из arange(TP), так что каждый ранг получает свой номер) и рассеивать через scatter_to_process_group_spmd. Перекомпиляция — «The capital of France is Paris.», SEQ_MATCH True, 77 мс prefill.

Почему это важно за пределами Inferentia

Эта история — не специфика AWS. Это универсальная архитектура проблем при переносе LLM на кастомное железо. Tensor parallelism, смешанные attention-механизмы, разделение параметров и буферов, зависимость от токенизатора — всё это встречается при портировании на TPU, Intel Gaudi, Habana, и даже при шардировании больших моделей на нескольких GPU через DeepSpeed или FSDP.

Конкретные цифры из этого кейса: 39 минут компиляции одного ранга для 31B, 182 ГБ пиковой памяти хоста на машине с 384 ГБ, 108 ГБ скомпилированных артефактов. Время prefill — 115 миллисекунд для dense, 77 миллисекунд для MoE. Всё это работает. После того, как исправлены все ловушки.

Для 26B-A4B дополнительные артефакты: Docker Hub xbill9/gemma4-optb-26b, HuggingFace xbill9/gemma-4-26B-A4B-it-inferentia2. Для 31B — только S3 с neff-файлами. Упаковка тоже подкинула проблем: при публикации сервер не мог загрузить модель из-за PytorchStreaming

Общий паттерн отладки: когда SEQ_MATCH = True, но вывод бессмысленный, баг не в ускорителе. Он в том, что разделяют CPU-референс и устройство — общий токенизатор, общие буферы, общие предположения о формате промпта. Unit-тесты проверяют численную эквивалентность компиляции, но не семантическую корректность вывода. Нужен отдельный тест на «модель отвечает осмысленно на простой вопрос» — и он должен использовать реальный чат-шаблон, а не хардкод маркеров, которые токенизатор разбивает на семь букв.

Часто задаваемые вопросы

Почему unit-тесты не поймали баг с токенизатором?

Потому что unit-тесты сравнивали модель с CPU-референсом, а референс использовал тот же сломанный токенизатор. Когда оба конца конвейера получают одинаковый некорректный вход, их выход совпадает идеально. Тест проверяет численную эквивалентность компиляции, но не семантическую корректность вывода. Нужен отдельный тест на «модель отвечает осмысленно на простой вопрос» — и он должен использовать реальный чат-шаблон, а не хардкод.

Зачем вообще портировать на Inferentia, если есть GPU?

Стоимость инференса. Inferentia2 предлагает значительно меньшую цену за токен для рабочих нагрузок, которые успешно портированы. Для команд, которые уже развернули инфраструктуру на AWS, Inferentia даёт 2–4× снижение cost-per-token по сравнению с эквивалентными GPU-инстансами. Цена — инженерное время на портирование и отлов ловушек вроде описанных выше. Для MoE-моделей добавляется расточительность FLOP: вычисляются все 128 экспертов, используются 8 — но это компромисс за трассируемость и корректность.

Что такое SPMD-ранг и почему он ломает MoE-модели?

SPMD (Single Program Multiple Data) — модель исполнения, где один скомпилированный граф запускается на всех рангах, но каждый ранг работает со своим куском данных. ModelBuilder компилирует один ранг и реплицирует граф. Если в графе зашита операция, зависящая от номера ранга (например, «взять срез 0:16 для экспертов»), то при компиляции ранга 0 это значение встраивается в граф как константа — и все 8 рангов начинают использовать срез ранга 0. Решение — SPMDRank-модуль, где номер ранга является параметром модели, загружаемым при старте, а не константой компиляции.

Как избежать тихой подмены буферов дефолтами?

Правило: всё, что зарегистрировано через register_buffer — послойные скаляры, некоторые нормализации, кастомные коэффициенты — требует явного копирования после загрузки шардированного чекпоинта. ModelBuilder и аналогичные системы перемещают только параметры. Буферы остаются в памяти модели с дефолтными значениями из конфига. Если слой масштабирует выход через layer_scalar, а буфер тихо дефолтнулся в 1.0 — шестьдесят слоёв ошибки сложатся в шум без единого сообщения об ошибке. Проверяйте: зарегистрированы ли у вашей модели буферы, и если да — добавьте ручное копирование в загрузчик весов.

Итог

Шесть моделей Gemma-4 (пять плотных + одна MoE), шесть принципиально разных ловушек при портировании на Inferentia2. От OOM при ручном трейсинге до тихой подмены буферов дефолтами, от невозможности шардировать 4 головы на 8 рангов до токенизатора, который превращает спецтокены в семь букв. От математически корректной «вычислить всех 128 экспертов и замаскировать ненужных в ноль» до SPMD-ранга, который встраивается как константа компиляции вместо рантайм-параметра. Общее правило: если ваш ускоритель совпадает с CPU-референсом токен-в-токен — радуйтесь ровно секунду, а потом проверьте, не сломан ли референс. Корректность компиляции ≠ корректность вывода. SEQ_MATCH True ≠ «The capital of France is Paris.»

← Все записи