Назад к блогу

Rufus-Air: как устроена пост-тренировка LLM на 106B параметров

Rufus-Air: как устроена пост-тренировка LLM на 106B параметров

Команда исследователей открыла рецепт пост-тренировки открытой модели GLM-4.5-Air-Base с 106 миллиардами параметров (12B активных), разбив весь процесс на восемь последовательных стадий — от SFT до RLHF с агентными сценариями. Статья подробно объясняет, как устроены роли на разных GPU, как генератор vLLM сосуществует с тренировочным процессом и какие подводные камни возникают при расчёте наград. Особый интерес представляет детерминированный анализ инфраструктурных решений: от формата датасетов и обработки `None`-наград до экстракторов ответов с поддержкой вложенных скобок.

Rufus-Air — открытый и воспроизводимый рецепт пост-тренировки на GLM-4.5-Air-Base (106B-A12B), собранный как последовательный конвейер из восьми стадий. Разбираем, что именно происходит на каждой стадии, как устроены роли и данные, как считается лосс и что ломается на границах.

Какие стадии покрывает рецепт

Конвейер состоит из восьми стадий: SFT, Reasoning RL, Coding RL, Instruction-Following RL, General Agent, Coding Agent, Search Agent и RLHF. SFT. Стадии Reasoning RL, Coding RL, Instruction-Following RL и RLHF — это RL. Награды на этих стадиях идут от жёстких проверяемых к более мягким сигналам на основе судей.

DPO в списке восьми стадий отсутствует.

Роли и данные между ними

Роли разделены по GPU. Тренировочный процесс с моделью и оптимизатором занимает один тренировочный GPU, а сервер vLLM — отдельный генерационный GPU; назначения не должны пересекаться. vLLM — внешний процесс, а не второй тренировочный процесс, поэтому размещение его на тренировочном GPU вызывает конкуренцию за одну и ту же память.

Rollout-воркер вызывает add_response_schema, чтобы разобрать завершения на content, <think>-рассуждения и блоки <tool_call>.

Судья интегрирован как пользовательская асинхронная reward-функция. Reward-функция вызывается с prompts, completions, completion_ids, плюс все остальные колонки датасета как именованные аргументы, а также trainer_state и log_metric.

Параметр only_label ограничивает оценку выборками с совпадающим значением колонки label; остальные получают None. Параметр reference_column задаёт колонку с эталонным ответом, читаемым из kwargs[reference_column] с откатом на "solution". Возврат None для выборки означает, что reward к ней не применяется, и TRL исключает её из вычисления через nansum/nanmean, а не оценивает как 0.0.

Формат промптов и датасетов

Датасеты бывают prompt-only и с эталонным ответом. В prompt-only формате строка содержит только prompt (список сообщений или строка), а если есть messages, загрузчик оставляет ходы до последнего ответа ассистента и отбрасывает эталонный ответ.

В GRPO при наличии messages столбец разбивается на prompt (все ходы до последнего ответа ассистента) и expected_output (содержимое последнего ответа ассистента или None, если ответа ассистента нет вовсе). Если есть только prompt, он нормализуется на месте, а остальные столбцы, включая solution, остаются нетронутыми.

Каждому датасету добавляется столбец label из ключа label: или путь по умолчанию, и datasets.concatenate_datasets выравнивает несовпадающие схемы, заполняя отсутствующие столбцы None. Награды используют эти столбцы: make_string_match_reward извлекает ответ из kwargs[answer_column] и сравнивает нормализованные строки.

Перед подачей в обучение применяется add_response_schema для разбора завершений на content, <think> и <tool_call>.

Генерация роллаутов

Промпт подаётся в Qwen2.5-0.5B-Instruct, затем происходит сэмплирование и получается 8 независимых роллаутов. Пять параметров — temperature, top_p, top_k, min_p, repetition_penalty — управляют сэмплированием студента.

Количество сэмплов на промпт различается по конфигурациям: в одном случае задано num_generations: 4, в другом получается 8 независимых роллаутов, а в третьем, в отличие от GRPO, num_generations отсутствует и делается одна генерация на промпт.

Детерминированный reward-функционал

make_string_match_reward извлекает финальный ответ из completion через extractor и такой же ответ из kwargs[answer_column] (при неудаче извлечения берётся сырая золотая строка), после чего сравнивает нормализованные строки и ставит 1.0/0.0.

Доступны четыре экстрактора. boxed берёт последний \boxed{...} со сбалансированным сканированием скобок (вложенные скобки работают). gsm8k берёт текст после последнего ####. last_number берёт последний числовой токен (необязательный знак, необязательная десятичная часть, запятые разрешены). full берёт весь completion без изменений.

Нормализация обеих сторон перед сравнением: strip, схлопывание внутренних пробелов, lowercase, удаление завершающей ., удаление запятых-разделителей тысяч внутри чисел, удаление окружающих $/\$, удаление окружающих {}.

Когда золотое значение отсутствует, пустое или не совпадает с меткой при заданном only_label, сэмпл не оценивается вовсе — он исключается из вычисления, а не получает оценку 0.0. То же происходит, если only_label задан, но ни в одном датасете смеси нет колонки label.

В отличие от accuracy_reward, который через math_verify считает 0.5 и 1/2 равными, string-match выбирают ради скорости, детерминизма и отсутствия зависимостей, когда золотые ответы уже в канонической извлекаемой форме. В playground V1 reward намеренно узкий: он проверяет, что последнее число в ответе равно ожидаемому, и ставит 1 или 0.

Фильтрация промптов по обучаемости

Группы формируются по наградам, полученным для каждого промпта:

  • все награды равны 1 — all-pass;
  • часть наград 0, часть 1 — learnable;
  • все награды 0 — all-fail.

Отбор выполняет модуль reward.filter_tasks, который выбирает полезные обучающие промпты, но не вычисляет преимущества и не обновляет параметры.

В GRPO награды нормализуются внутри каждой группы из num_generations сэмплов, и если все сэмплы получают одинаковую награду, каждый токен получает нулевое преимущество и группа не даёт градиента. Синхронный GRPO борется с этим через dynamic_sampling. AsyncGRPO динамическое сэмплирование отклоняет.

При resample информативные группы сохраняются целиком, а новая партия берётся из отдельного сэмплера с repeat_count=1 и seed+1, генерируется, оценивается, и её информативные группы добавляются, пока локальная партия не заполнится или не будет достигнут dynamic_sampling_max_rounds. Режим mask применяется в _prepare_inputs к микропартии TRL только когда каждая строка в ней «мёртвая», поскольку прямоугольный тензор нельзя укоротить для одних строк и не для других.

Групповые преимущества и GRPO-лосс

Group-relative advantage вычисляется так: награда выше среднего по группе даёт положительное преимущество и поощряет поведение, ниже — отрицательное и подавляет, а одинаковые награды дают нулевые преимущества без контраста.

При нулевой дисперсии, когда все награды в группе одинаковы, все преимущества становятся нулевыми, и группа не даёт градиента. Это проявляется как loss: 0 / grad_norm: 0, поскольку преимущество GRPO — это отклонение награды от среднего по своей группе, и группа с одинаковыми наградами не даёт градиента вовсе.

Клиппированный суррогатный лосс GRPO строится из отношения вероятностей old/new: берутся старые и новые логарифмы вероятностей, из них получается отношение вероятностей, которое умножается на advantage, затем «слишком большие благоприятные изменения» обрезаются (clip), после чего результат усредняется и берётся со знаком минус, чтобы получить минимизируемый лосс.

Вклад токена обнуляется, когда строка «мертва». Deadness определяется по группе при скоринге батча, а не по отдельной строке микробатча, поскольку строка ровно на среднем своей группы имеет advantage 0, хотя группа информативна, и её маскирование отбросило бы реальный градиент. Режим mask оставляет градиент нетронутым только если мёртвая строка действительно ничего не вносит, что выполняется при beta: 0, отсутствии entropy bonus и router auxiliary loss. Режим resample намеренно меняет градиент, заменяя строки с нулевым advantage на информативные, поэтому прогон с ним не воспроизводим относительно прогона без него.

Обновление политики и LoRA

В V7 исходный чекпоинт Qwen замораживается, к проекциям внимания прикрепляются LoRA-адаптеры ранга 4, и обновляются только параметры адаптеров с помощью supervised-лосса по токенам ответа. Обученное поведение хранится в адаптере LoRA, поэтому загрузка одной базовой модели Qwen его не включает.

Тренер накладывает два ограничения на конфигурацию адаптеров: lm_head не должен входить в target_modules, потому что лосс читает lm_head.weight напрямую и адаптер там никогда не обучался бы; и не поддерживаются prompt-learning методы (PromptTuning / PrefixTuning / P-Tuning) — вместо них следует использовать обычный LoRA.

Категориальный KL-штраф

Категориальный KL-штраф задаётся параметром beta, который интерполирует саму функцию потерь между разными дивергенциями:

beta: 1.0   # reverse KL, KL(student || teacher) <- on-policy distillation (default)
beta: 0.0   # forward KL, KL(teacher || student) <- mass-covering
beta: 0.5   # Jensen-Shannon divergence

reverse KL является mode-seeking: студент фиксируется на одном поведении учителя вместо размывания нескольких вместе, тогда как forward KL заставляет студента покрывать все моды учителя. При beta: 0.5 получается Jensen-Shannon divergence. По умолчанию рекомендуется оставлять beta равным 1.0, если только нет специального намерения заставить студента покрыть все моды учителя.

Этот beta не следует путать с beta из GRPO: там это коэффициент KL-штрафа против референсной модели, а здесь он интерполирует саму функцию потерь, и референсной модели нет. Потери считаются по-токенно: loss — это средняя по-токенная дивергенция в натах, и loss ≈ 0 означает, что студент воспроизводит распределение учителя на собственных роллаутах.

Сохранение, перезагрузка и возобновление

При сохранении чекпоинта функция _save_checkpoint записывает файл rollout_state.json рядом с чекпоинтом модели. При возобновлении этот файл используется, чтобы переустановить курсор набора данных rollout-воркера. Resume работает корректно: загрузка checkpoint-6 из двенадцатишагового запуска возобновилась на глобальном шаге 6 и завершила оставшиеся шаги.

Однако поле prompt_index в rollout_state.json — это нижняя граница (low-water mark), а не счётчик завершённых промптов: TRL вычисляет его как наименьший индекс группы, отсутствующий в множестве групп, дошедших до модели. Это плохо взаимодействует с отбрасыванием устаревших сэмплов: путь stale-discard не записывает группу отброшенного сэмпла, поэтому если все сэмплы группы отброшены, группа никогда не попадает в множество завершённых, и отсутствующая группа с малым номером может бессрочно удерживать prompt_index на месте, даже когда обучение уходит далеко вперёд.

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

Асинхронная топология и неподдерживаемые конфигурации

GPU выбираются через независимые настройки CUDA_VISIBLE_DEVICES, и их назначения не должны пересекаться. AsyncGRPO явно не поддерживает DeepSpeed, FSDP, несколько тренировочных процессов, запуски оценки, динамическую выборку, квантованные модели, PEFT и модель вознаграждения, размещённую на локальном GPU. Функция validate_async_args в dna_factory/async_grpo.py отклоняет каждую из этих конфигураций до загрузки модели или датасета.

Общая предпочтительность DeepSpeed над FSDP в репозитории относится только к режимам обучения, поддерживающим распределённое выполнение, и не подразумевает поддержку DeepSpeed для AsyncGRPO.

Ограничение устаревания роллаутов

Устаревание. RolloutQueueDataset отбрасывает любой образец, чьё устаревание превышает max_staleness, и вместо него берёт следующий доступный образец.

Отслеживаются метрики sample/staleness_mean, sample/staleness_max, sample/dropped_stale_total, perf/rollout_wait_s, sample/rollout_queue_size. При max_inflight_tasks: -1 значение разрешается по формуле:

max_staleness x per_device_train_batch_size x gradient_accumulation_steps x num_processes

Формула предполагает, что образец в полёте будет потреблён в пределах настроенного окна устаревания. Длинные completion'ы могут нарушить это предположение: генерация обгоняет потребление, и избыточные образцы устаревают до того, как обучение до них дойдёт. Отброшенные устаревшие образцы — неотъемлемая цена разделения генерации и оптимизации, а их частота зависит от баланса между скоростью генерации и скоростью потребления при обучении.

token_budget и length-aware batching

token_budget ограничивает число непаддинговых токенов, упакованных в одну строку одного forward pass: короткие последовательности могут делить строку, а длинные занимают строку целиком, что ограничивает пиковую память независимо от числа сэмплов за шаг. Если сэмпл сам превышает бюджет, он не помещается ни в одну строку, и TRL отбрасывает его с предупреждением и продолжает обучение.

По умолчанию в configs/_defaults-AsyncGRPO.yaml задано token_budget: null. При null TRL читает max_model_len vLLM-сервера при старте обучения, что выполняет инвариант контекста по построению и избегает устаревшего жёсткого значения.

Слишком маленький бюджет вредит в первую очередь пропускной способности, а не корректности. При token_budget: 4096 и 2048-токенных завершениях каждая строка держала только один сэмпл, шаг потреблял четыре сэмпла и покрывал лишь половину группы генерации, за двенадцать шагов обучено только семь групп. После увеличения контекста сервера до 16384 и оставления token_budget null примерно два сэмпла помещались в строку, и те же двенадцать шагов обучили одиннадцать групп.

num_generations по-прежнему определяет группы для нормализации награды, но эти группы могут охватывать обновления оптимизатора, так как батчер не сохраняет границы групп генерации, поэтому пошаговые reward_std и grad_norm не должны соответствовать один-к-одному.

Пограничные случаи

Zero-advantage группы

При одинаковой награде всех сэмплов группы каждый токен получает нулевой advantage, и группа не даёт градиента. Синхронный GRPO обрабатывает это через dynamic_sampling режимами mask или resample, а AsyncGRPO динамическое сэмплирование отклоняет.

mask применяется в _prepare_inputs на микробатче, только когда все строки в нём мёртвые, и оставляет градиент нетронутым лишь при beta: 0, без entropy bonus и без router auxiliary loss. resample намеренно меняет градиент, заменяя нулевые строки информативными, поэтому прогон с ним не воспроизводим относительно прогона без него; он пересчитывает нормализатор inputs["num_items_in_batch"] для пополненного батча.

Метрика dyn/dead_frac — доля локальных строк в мёртвой группе после любого пополнения, а dyn/refilled равен 1, когда пополнение заполнило батч, и 0 при откате. dyn/dead_frac не совпадает с frac_reward_zero_std, который TRL измеряет до фильтрации, по промптам и собранным по ранкам.

Превышение длины completion над контекстом

Когда запрошенная длина вывода превышает доступный контекст, vLLM не обрезает запрос, а возвращает HTTP 400 с сообщением о максимальной длине контекста модели. TRL повторяет такой запрос тридцать раз, прежде чем дочерний процесс rollout-worker завершается, и на верхнем уровне видна только ошибка RuntimeError: AsyncRolloutWorker child exited during init с вложенной ClientResponseError: 400, message='Bad Request'.

Чтобы этого избежать, должны одновременно выполняться два условия:

--max-model-len >= longest prompt + max_completion_length
token_budget >= longest prompt + max_completion_length

Функция validate_completion_length_against_server в dna_factory/async_grpo.py проверяет это при старте: читает max_model_len сервера через GET /v1/models тем же способом, что и VLLMClient.get_max_model_len в TRL, затем поднимает ошибку, если max_completion_length больше или равен max_model_len, предупреждает, если непустой token_budget меньше max_model_len, и логирует оставшийся запас для промпта. setup_async_training_args выполняет эту проверку до загрузки токенизатора, датасета или модели, превращая много минутную ошибку в ошибку на старте.

Проверка сервера выполняется по мере возможности: если сервер ещё недоступен после нескольких коротких повторов, она логирует, что проверка пропущена, и продолжает, оставляя обработку неготового сервера функции TRL wait_for_server_ready.

Оценка чекпоинтов и метрики

Колбэк checkpoint-eval включается флагом eval_on_checkpoint (по умолчанию false) и запускается в общем runner для sft.py, dpo.py, grpo.py, distill.py. При eval_baseline: true (по умолчанию) on_train_begin запускает первый eval до первого шага оптимизатора против модели, с которой начинается обучение: vllm serve <model_name_or_path> --served-model-name <tag>-checkpoint-0, логи идут в <output_dir>/eval_logs/step-0/.

Каждая задача из набора по умолчанию становится одной метрикой eval/<task>, например eval/mmlu_pro, eval/gpqa_diamond, eval/kmmlu_pro. Результаты логируются в W&B против выделенной оси eval/step, объявленной через define_metric, и тот же шаг пишется как train/global_step:

run.log({"eval/step": 1776, "train/global_step": 1776, "eval/kmmlu_pro": 0.61})

Для кривых report_to должен включать wandb; без этого оценки всё равно появляются в логе обучения, а при старте выводится предупреждение.

Гарантии eval-колбэка по ранкам и размещению GPU

Eval-колбэк гейтится на state.is_world_process_zero: все действия выполняет только нулевой ранг, остальные ранги выходят из on_save без изменений, и никакой коллективной операции (NCCL) не публикуется, поскольку согласовывать между рангами нечего.

Чтобы ротация чекпоинтов не удалила тот, что сейчас оценивается, on_save сразу делает снимок в <output_dir>/_eval_staging/checkpoint-<N> (веса, токенизатор и конфиги, без optimizer.pt), используя os.link — без копирования и диска, а inode переживают rmtree ротации; при отказе линковки (другая файловая система) выполняется реальное копирование с предупреждением. Одновременно идёт только один eval: пришедший во время текущего чекпоинт пропускается с предупреждением, а не ставится в очередь.

Сервер vLLM запускается в отдельной сессии и снимается через os.killpg на нормальном пути, при исключении, в on_train_end и из atexit-хука; утечка возможна только при SIGKILL самого тренера.

GPU для eval-сервера должны быть отдельными: пересечение eval_devices с CUDA_VISIBLE_DEVICES тренировочного процесса даёт громкое предупреждение при старте, а пустой eval_devices при включённом переключателе — ошибку старта. По умолчанию это две GPU 0,1 с --data-parallel-size 2, и eval_devices читается как физические id устройств, задаваемые в окружении сервера дословно, независимо от CUDA_VISIBLE_DEVICES тренера.

Компромиссы: почему reverse KL и почему dynamic sampling меняет градиент

По умолчанию выбран reverse KL (beta=1.0), потому что он mode-seeking: студент фиксируется на одном поведении учителя, а не размывает несколько вместе, и такой лосс невозможно «взломать» — низкий KL может означать только то, что студент воспроизводит поведение учителя. Forward KL (beta=0.0) описан как mass-covering и применяется, только если специально нужно, чтобы студент покрывал все моды учителя.

On-policy подход (студент сам сэмплирует свои завершения) устраняет exposure bias: off-policy дистилляция — обычный SFT на выходах учителя — показывает студенту только те состояния, которые посещает учитель, и на инференсе студент накапливает ошибки из состояний, на которых не обучался.

Компромисс по сравнению с RL: дистилляция даёт по-токеновый сигнал учителя, тогда как RL сворачивает всю траекторию в одну скалярную награду, что даёт выигрыш на порядок по вычислительной эффективности. Промпты можно безопаснее переиспользовать между эпохами, чем в RL, поскольку по-токеновый сигнал учителя продолжает супервизировать даже на уже виденном промпте.

Dynamic sampling меняет градиент намеренно в режиме resample, потому что замена строк с нулевым преимуществом на информативные и есть цель этого режима. Режим mask, наоборот, оставляет градиент нетронутым, но только когда мёртвая строка действительно ничего не вносит — при beta: 0, без entropy bonus и без router auxiliary loss; все три условия проверяются при инициализации и о них предупреждают, прогон продолжается.

Нормализатор DAPO — это inputs["num_items_in_batch"], скаляр, фиксируемый при оценке батча, поэтому усечение мёртвых строк не перемасштабирует выживших, а resample пересчитывает его для пополненного батча. Из компонентов DAPO в TRL есть loss_type="dapo", epsilon_high, mask_truncated_completions, но не dynamic sampling; сравнимые реализации — verl FilterGroupsConfig и ms-swift GRPOTrainer._dynamic_sampling.

resample не поддерживается со streaming datasets и падает при инициализации, а оба режима не поддерживаются с multimodal или token-type входами (pixel_values, token_type_ids и т.п.), чья вторая ось не является ни промптом, ни завершением, — падает на первом батче.

Воспроизведение и запуск

Минимальный запуск состоит из двух процессов на разных GPU. Сервер генерации запускается с CUDA_VISIBLE_DEVICES=1 и обязательной переменной VLLM_SERVER_DEV_MODE=1; последняя нужна потому, что маршруты /pause, /resume, /init_weight_transfer_engine, /start_weight_update, /update_weights, /finish_weight_update, /get_world_size, /server_info регистрируются только в dev-режиме. Сервер обслуживает модель dnotitia/Qwen3-0.6B с режимом логвероятностей processed_logprobs, передачей весов через бэкенд nccl, типом bfloat16 и максимальной длиной контекста 16384.

Обучение запускается в отдельном терминале на GPU с CUDA_VISIBLE_DEVICES=0:

CUDA_VISIBLE_DEVICES=0 python grpo.py \
  --config configs/GRPO/qwen3-0.6B-async.yaml \
  --grpo_execution async \
  --max_completion_length 8192

Асинхронный режим включается либо в YAML (grpo_execution: async), либо флагом --grpo_execution async, при этом CLI переопределяет YAML.

Для judge-сервера настраивается только соединение через переменные окружения: JUDGE_BASE_URL (по умолчанию http://localhost:8001/v1), JUDGE_MODEL (если не задана — используется первая модель из списка сервера) и JUDGE_API_KEY (по умолчанию EMPTY). Остальные параметры judge зафиксированы в generative.py как константы: конкурентность 16, таймаут 120s, 2 повтора с экспоненциальной задержкой 1s и 2s, 2048 максимум генерируемых токенов, порог бинаризации 7 и колонка "solution".

Что из этого следует на практике

  • Разделение тренировочного и генерационного GPU обязательно: vLLM — внешний процесс, и совмещение его с обучением на одном GPU приводит к конкуренции за память.
  • token_budget: null предпочтительнее жёсткого значения: TRL читает max_model_len сервера при старте, и инвариант контекста выполняется по построению. Слишком маленький бюджет бьёт по пропускной способности, а не по корректности.
  • Zero-advantage группы не дают градиента. AsyncGRPO не имеет механизма их восстановления, поэтому доля мёртвых групп напрямую влияет на эффективность обучения.
  • resample делает прогон невоспроизводимым относительно прогона без него, поскольку намеренно меняет градиент.
  • prompt_index в rollout_state.json — нижняя граница, а не счётчик. При отбрасывании устаревших сэмплов resume может заново обработать уже пройденные группы: это дублирование работы, а не потеря данных.
  • max_completion_length и token_budget должны покрывать сумму самого длинного промпта и максимальной длины завершения, иначе vLLM вернёт HTTP 400, а TRL тридцать раз повторит запрос, прежде чем rollout-воркер завершится с неинформативной ошибкой верхнего уровня.
  • Eval-сервер требует отдельных GPU: пересечение eval_devices с CUDA_VISIBLE_DEVICES тренера даёт предупреждение при старте, а пустой eval_devices при включённом переключателе — ошибку.

Где смотреть в коде

Источники

Похожее