Direct Preference Optimization (DPO) дозволяє розробникам донавчати великі мовні моделі без необхідності тренування окремої моделі винагороди (reward model) або запуску циклу навчання з підкріпленням (RL), що суттєво знижує як витрати на обчислення, так і нестабільність, яка часто притаманна традиційним конвеєрам RL-з-відгуками-людей (RLHF).
Чому RL-з-відгуками-людей (RLHF) здається занадто складним
Стандартний підхід RL-з-відгуками-людей (RLHF) складається з трьох етапів. По-перше, базова модель донавчається на кураторському наборі даних. По-друге, модель винагороди вчиться передбачати людські вподобання між парами відповідей. Нарешті, фахівці запускають Proximal Policy Optimization (PPO) — класичний алгоритм RL — щоб спрямувати стратегію (policy) до вищих передбачених винагород, залишаючись при цьому близькою до оригінальної моделі.
Така конфігурація з трьох моделей одночасно займає пам'ять і вимагає дорогого циклу RL, який здійснює вибірку (sampling) нового тексту під час кожного оновлення. Команди часто спостерігають, як стратегія вчиться «обманювати» систему винагороди, створюючи відповіді, які мають високі бали за проксі-метриками, але не відповідають бажаній якості. У результаті виходить конвеєр, який є дорогим, крихким і важким для масштабування.
Алгебраїчний «короткий шлях» DPO
DPO повністю пропускає етап створення моделі винагороди. Ключове спостереження полягає в тому, що ціль RLHF — максимізація очікуваної винагороди при одночасному штрафуванні за відхилення від референсної моделі — має вираз у замкненій формі. Перегрупувавши математичні формули, можна побачити, що винагорода для будь-якого токена стає різницею між логарифмічною ймовірністю, призначеною стратегією, та логарифмічною ймовірністю, призначеною референсною моделлю.
На практиці це означає, що «винагорода» міститься безпосередньо в самій стратегії. Навчання зводиться до єдиної функції втрат класифікації на парах вподобань: маючи обрану відповідь і відхилену, модель стимулюється призначати вищу ймовірність обраному тексту. Жодної вибірки, жодних оновлень PPO, жодних додаткових моделей для зберігання.
Як виглядає нова функція втрат
Функція втрат порівнює логарифмічну ймовірність кращої відповіді за поточною стратегією з логарифмічною ймовірністю за референсною моделлю, масштабуючи це значення за допомогою гіперпараметра β, подібного до температури. Високе значення β змушує стратегію залишатися близько до референсної моделі, зберігаючи плинність мови та безпеку. Низьке значення β дозволяє стратегії відхилятися далі, посилюючи її вподобання щодо обраної відповіді.
Переваги, важливі для розробників
- Відсутність моделі винагороди — усуває потребу в зборі додаткових відгуків людей для окремого предиктора.
- Відсутність вибірки під час навчання — модель ніколи не генерує новий текст для обчислення градієнтів, що суттєво скорочує час роботи GPU.
- Стабільність — стандартна функція втрат binary-cross-entropy замінює градієнти RL з високою дисперсією, які часто спричиняють розбіжність.
- Ефективність — достатньо одного кроку градієнта для кожної пари вподобань; навчання сходиться за набагато меншу кількість епох, ніж у PPO.
Перші експерименти показують, що DPO зрівнюється з PPO або перевершує його за продуктивністю на бенчмарках наборів даних вподобань, використовуючи при цьому лише частку обчислювального бюджету. Ця перевага в вартості пояснює, чому багато open-source проєктів уже прийняли DPO або його близький варіант як метод вирівнювання (alignment) за замовчуванням.
Компроміси
DPO працює на фіксованому наборі пар вподобань. Оскільки під час навчання вона ніколи не здійснює вибірку нових варіантів завершення тексту, вона не може досліджувати простір відповідей, яких не було в оригінальних даних. Натомість онлайн-запуск PPO може виявляти нові, більш вигідні поведінки шляхом постійного зондування моделі.
Якщо β встановлено занадто низьким або навчання триває занадто багато кроків, стратегія може відхилитися від референсної моделі настільки, що втратить плинність мови або почне створювати небажані артефакти.
Підсумок
Direct Preference Optimization замінює важкий для RL стек RLHF із трьох моделей на єдину стабільну функцію втрат, яка навчається безпосередньо на парах людських вподобань. Результатом є дешевший і передбачуваніший шлях до вирівнювання мовних моделей — за умови, що навчальні дані охоплюють необхідні вам поведінки, а ви контролюєте параметр відхилення.
