Перейти к содержанию
Публикация AiManual

RLOO: лёгкая и быстрая альтернатива PPO для онлайн RLHF-обучения

RLOO требует на 50–70% меньше видеопамяти и работает в 2–3 раза быстрее PPO при сопоставимом качестве. Разбираем архитектуру, результаты на Pythia 1B и 6.9B, на

Коротко

Что будет в материале

  1. 01

    Что такое RLOO и почему он появился?

  2. 02

    Ключевые отличия RLOO от PPO

  3. 03

    Эксперименты: RLOO против PPO на Pythia 1B и 6.9B

  4. 04

    Практическое использование RLOO в TRL

RLOO (REINFORCE Leave One-Out) - алгоритм онлайн RLHF-обучения, который Cohere представил как замену PPO для задач выравнивания языковых моделей. Он требует на 50–70% меньше видеопамяти, работает в 2–3 раза быстрее и показывает сопоставимое качество ответов. В экспериментах на Pythia 1B и 6.9B RLOO достиг win rate 78.7% против 77.9% у PPO. Алгоритм уже интегрирован в библиотеку TRL, поэтому его можно запустить без самостоятельной реализации.

Главное отличие RLOO от PPO - отказ от value-функции и клиппинга. RLOO моделирует всю сгенерированную последовательность как единое действие, использует REINFORCE-лосс и вычисляет бейзлайн методом leave-one-out по батчу. Это сокращает число копий модели с четырёх до трёх и упрощает вычислительный граф.

Материал разбирает архитектуру RLOO, результаты сравнения с PPO, практическую настройку в TRL и подводный камень: численную нестабильность logprob при bf16, которая зануляет градиенты у 20–40% батча.

Что такое RLOO и почему он появился?

RLOO расшифровывается как REINFORCE Leave One-Out. Алгоритм относится к семейству policy gradient методов и предназначен для онлайн RLHF, когда модель генерирует ответы, получает награду от reward-модели и обновляет свои веса в реальном времени. Cohere представила RLOO как практичный ответ на главную проблему PPO - высокую стоимость обучения.

Проблемы PPO в RLHF

PPO требует четыре копии модели в памяти: policy, reference, value и reward. Каждая копия занимает видеопамять, а при больших LLM это быстро упирается в лимиты GPU. Value-функция добавляет отдельную голову, которую нужно обучать, а клиппинг усложняет лосс и увеличивает число гиперпараметров. На практике PPO для RLHF часто оказывается медленным и нестабильным, особенно на ограниченном железе. Подробнее о том, как PPO применяется в RLHF и какие нестабильности возникают, разобрано в материале про StackLLaMA.

Идея RLOO: моделирование генерации как единого действия

RLOO рассматривает всю сгенерированную последовательность как одно действие. Вместо того чтобы оценивать каждый токен отдельно, алгоритм берёт полный ответ, получает за него скалярную награду и обновляет вероятность генерации всей последовательности. Это позволяет применить REINFORCE-лосс без value-функции и без клиппинга. Бейзлайн для снижения дисперсии градиента вычисляется методом leave-one-out по батчу: для каждого сэмпла бейзлайном служит средняя награда остальных сэмплов в том же батче.

Ключевые отличия RLOO от PPO

Сравнение удобно свести к четырём параметрам: число копий модели, тип лосса, способ подсчёта бейзлайна и вычислительные затраты.

  • Копии модели: RLOO держит три копии (policy, reference, reward), PPO - четыре (плюс value).
  • Лосс: RLOO использует REINFORCE, PPO - PPO-clip.
  • Бейзлайн: RLOO считает leave-one-out по батчу, PPO обучает value-функцию.
  • Скорость и память: RLOO на 50–70% экономичнее по видеопамяти и в 2–3 раза быстрее.

REINFORCE-лосс вместо PPO-клиппинга

REINFORCE обновляет политику по формуле: градиент логарифма вероятности действия, умноженный на преимущество. В RLOO преимущество - это разница между наградой сэмпла и бейзлайном. Нет клиппинга, нет value-функции, нет отдельной оптимизации критика. Это сокращает число гиперпараметров и упрощает отладку. Плата за простоту - потенциально более высокая дисперсия градиента, которую RLOO компенсирует leave-one-out бейзлайном.

Бейзлайн на основе leave-one-out

Для каждого сэмпла в батче RLOO вычисляет бейзлайн как среднюю награду всех остальных сэмплов. Если в батче K генераций на один промпт, то для i-го сэмпла бейзлайн равен среднему по K-1 остальным наградам. Такой подход снижает дисперсию градиента без дополнительных параметров и без отдельной модели. Число генераций на промпт напрямую влияет на стабильность бейзлайна: чем больше K, тем надёжнее оценка.

Эксперименты: RLOO против PPO на Pythia 1B и 6.9B

Сравнение проводилось на моделях Pythia 1B и Pythia 6.9B в задаче RLHF. Оценивалось качество ответов через win rate и вычислительная эффективность.

Метрики и win rate

Win rate показывает, как часто ответы обученной модели предпочитают эталонным ответам при сравнении. RLOO достиг 78.7%, PPO - 77.9%. Разница в 0.8 процентного пункта находится в пределах шума, поэтому корректный вывод: качество сопоставимо, RLOO не уступает PPO.

Скорость и потребление памяти

RLOO требует на 50–70% меньше видеопамяти за счёт отказа от value-модели: три копии вместо четырёх. Обучение идёт в 2–3 раза быстрее. Для команд, которые дообучают LLM на ограниченном числе GPU, это решающий аргумент. Связка PEFT, квантизации и RLHF на одной 24 ГБ GPU разобрана в статье про RLHF для 20B-моделей.

Практическое использование RLOO в TRL

RLOO интегрирован в библиотеку TRL, поэтому запуск сводится к выбору конфигурации. Базовый пример:

from trl import RLOOConfig, RLOOTrainer

config = RLOOConfig(
    learning_rate=1e-5,
    per_device_train_batch_size=4,
    num_ppo_epochs=1,
    num_mini_batches=1,
    rloo_k=4,
)
trainer = RLOOTrainer(
    config=config,
    model=model,
    ref_model=ref_model,
    reward_model=reward_model,
    train_dataset=dataset,
)
trainer.train()

Настройка и гиперпараметры

Ключевой параметр - rloo_k, число генераций на промпт. При K=2 бейзлайн считается по одному соседнему сэмплу, что даёт высокую дисперсию. Рекомендуется K от 4 до 8. Learning rate для RLHF обычно ниже, чем для SFT: 1e-5 или 5e-6. Batch size стоит подбирать так, чтобы в батче было не менее 8–16 генераций суммарно. Мониторить нужно среднюю награду, KL-дивергенцию от reference-модели и долю нулевых градиентов.

Подводные камни: численная нестабильность logprob при bf16

bf16 экономит память, но имеет меньшую точность, чем fp32. При вычислении logprob для длинных последовательностей значения могут уходить в -inf или nan. Это приводит к нулевым градиентам: модель перестаёт обновляться на части батча.

Причины и проявления

Проблема возникает из-за накопления ошибок округления при перемножении вероятностей токенов. В bf16 мантисса короче, поэтому logprob длинного ответа теряет точность. Симптомы: нестабильный loss, медленная сходимость, зануление градиентов у 20–40% батча. Обнаружить можно, если логировать долю сэмплов с нулевым или nan-градиентом.

Рекомендации по устранению

Первый шаг - вычислять logprob в fp32, даже если веса модели в bf16. Второй - добавлять epsilon к вероятностям перед логарифмированием. Третий - применять gradient clipping, чтобы ограничить влияние выбросов. Четвёртый - мониторить долю нулевых градиентов и при росте выше 10–15% переключать вычисление лосса на fp32. Эти меры убирают проблему в большинстве конфигураций.

Заключение: стоит ли переходить на RLOO?

RLOO - практичная альтернатива PPO для онлайн RLHF, когда ресурсы ограничены. Качество сопоставимо: 78.7% против 77.9% win rate. Скорость выше в 2–3 раза, видеопамяти нужно на 50–70% меньше. Алгоритм интегрирован в TRL и запускается с минимальными изменениями кода.

Численная нестабильность logprob при bf16 - реальный риск, но он решается вычислением лосса в fp32 и мониторингом градиентов. Для команд, которые уже используют PPO и упираются в память GPU, переход на RLOO выглядит оправданным. Для тех, кто только планирует RLHF, RLOO стоит рассмотреть как стартовый вариант из-за простоты и меньших требований к железу. О других подходах к оптимизации обучения LLM можно прочитать в разборе 20B Looping Model и материале про метод PoLar.

Подписаться на канал