В статье я попробую осветить подход к обучению нейронной сети игре в крестики-нолики с помощью методов обучения с подкреплением (Reinforcement Learning или RL). Мы разберем основные идеи TD Learning и Q-Learning, посмотрим, как сеть постепенно учится принимать все более сильные решения.
Статья не претендует на исчерпывающее изложение темы. Цель скорее в том, чтобы дать интуитивное понимание ключевых идей. Тем не менее надеюсь что материал покажется вам интересным и будет полезен.
Итак, все мы умеем играть в крестики-нолики - это очень простая игра, которая кроме того хорошо подходит для введения в тему самообучения(без учителя). Мы рассмотрим несколько методов - Backward TD(0), Batch (Forward) TD(0), Monte Carlo, Online TD(0) и Q-learning.*
*Здесь TD - Temporal Difference, а TD(0) - обозначает bootstrap на 1 шаг("подтягивание" к оценке соседнего состояния - об этом ниже).
Исходный код методов: исходный код на github
1.Краткая теория
1.1 Игра как задача обучения с подкреплением
Партию в крестики-нолики можно представить как марковский процесс принятия решений, где состояние - это текущее состояние игрового поля(доски), а действие - выбор свободной клетки для хода, вознаграждение равно нулю на протяжении всей партии и становится ненулевым только в конце (+1 победа X, −1 победа O, 0 ничья). Такая структура с редким, отложенным вознаграждением типична для настольных игр и делает задачу в общем случае очень не простой (хотя в случае крестиков-ноликов это не совсем так, но это же учебный пример), где сеть должна научиться связывать ранние ходы с далёким по времени исходом партии.
Цель обучения в том, чтоб найти функцию ценности, которая позволяет сети оценивать, насколько хороша та или иная позиция (или ход), и на основе этой оценки выбирать наилучшие действия.
1.2 Уравнение Беллмана
В основе почти всех методов RL лежит уравнение Беллмана. Функция ценности состояния по определению - это ожидаемая отдача (return) при нахождении в этом состоянии:
V(s) = E[Gₜ | Sₜ = s]
где Gₜ - суммарное (возможно, дисконтированное) вознаграждение до конца эпизода. Уравнение Беллмана переписывает это ожидание рекурсивно через ценность следующего состояния:
V(s) = E[Rₜ₊₁ + γV(Sₜ₊₁) | Sₜ = s]
Это принципиальный шаг: вместо того чтобы «ждать» до конца партии, можно оценивать состояния через оценку соседних состояний. Такой приём называется бутстрэпингом (bootstrapping) и лежит в основе всех TD-методов.
Для задач управления (control), где нужно не просто оценить политику, а найти оптимальную, используется уравнение оптимальности Беллмана уже не для V(s), а для функции ценности действия Q(s, a):
Q*(s, a) = E[Rₜ₊₁ + γ · maxₐ' Q*(Sₜ₊₁, a′)]
Это уравнение лежит в основе Q-learning.
1.3 Подход V(s) против подхода Q(s, a)
Есть два принципиально разных способа представить знания сети об игре.
V(s) - ценность состояния. Сеть оценивает, насколько выгодна позиция сама по себе, без привязки к конкретному ходу. Чтобы выбрать действие, сети приходится перебрать все возможные ходы, мысленно применить каждый к доске и оценить получившиеся состояния(afterstate-ы), т.е. сети нужна модель среды (а точнее - знание, как действие меняет состояние).
Q(s, a) - ценность действия в состоянии. Сеть сразу оценивает пару позиция + ход, и оптимальное действие находится напрямую через argmaxₐ Q(s, a), без необходимости заглядывать вперёд и симулировать переходы. Это делает Q-learning методом model-free control в чистом виде, когда сети не нужно знать правила игры для выбора хода - только для их совершения.
1.4 TD-ошибка и bootstraping
Уравнение Беллмана связывает V(s) и V(s'), но само по себе не даёт правила обучения. Оценки V сетью в начале обучения всегда крайне неточны. TD-методы представляют собой алгоритмы, сводящие уравнение Беллмана в правило обновления весов через понятие TD-ошибки (temporal difference error):
δₜ = target − V(Sₜ), где target = Rₜ₊₁ + γV(Sₜ₊₁)
TD-ошибка это просто разница между тем, что сеть предсказывала для Sₜ до хода, и тем, что получилось, когда стало известно следующее состояние Sₜ₊₁ (или, для терминального состояния, реальный исход партии). Если δₜ = 0 - сеть уже была права, обновлять нечего. Если δₜ ≠ 0 - оценка Sₜ подтягивается в сторону target-а на шаг обучения α:
V(Sₜ) ← V(Sₜ) + α · δₜ
Для Q-learning формула TD-ошибки аналогична, только с максимумом по действиям и с поправкой на смену игрока:
δₜ = target − Q(Sₜ, Aₜ), где target = −maxₐ' Q(Sₜ₊₁, a′) (или реальный исход, если Sₜ₊₁ терминально)
1.5 Почему все эти методы вообще работают
Все рассматриваемые методы - по-сути разные способы приблизить истинную функцию ценности через многократные партии self-play(т.е. игру сети с собой):
TD-методы (Backward TD(0), Batch TD(0), Online TD(0), Q-learning) стохастически аппроксимируют уравнение Беллмана, т.е вместо точного вычисления ожидания берется одна выборка перехода, а вместо истинной
V(s')подставляется текущая, ещё не идеальная оценка сети. Это классическая стохастическая аппроксимация (на подобии метода Роббинса–Монро): при достаточно малом шаге обучения и достаточном числе итераций оценка сходится к истинной функции ценности.Monte Carlo не использует рекурсию уравнения Беллмана вообще, а аппроксимирует напрямую определение
V(s) = E[Gₜ], используя реально дошедший до конца партии результат как несмещённую (но зашумлённую) выборку из этого ожидания.
Компромисс между этими двумя подходами - смещение против дисперсии (bias–variance tradeoff):
бутстрэпинг (TD) даёт смещённую, но низкодисперсную оценку и обучение эффективнее использует данные, но подвержено ошибкам из-за неточности текущей сети.
отсутствие бутстрэпинга (Monte Carlo) даёт несмещённую, но высокодисперсную оценку - точнее в среднем, но каждая отдельная партия сильнее "шумит".
1.6 Компромисс между исследованием и использованием: ε-greedy
Важная составляющая всех рассмотренных в статье методов это то, как именно сеть выбирает ходы во время обучения. Это не относится к тому, как обновляется функция ценности (TD, MC или Q-learning), а к тому, какие партии сеть вообще успевает сыграть, чтобы было что обновлять - так называемая дилемма exploration vs exploitation (исследование против использования).
Если сеть всегда жадно играет лучший ход по текущей (ещё не обученной) оценке, есть риск застрять в локально "уверенных", но объективно слабых стратегиях - сеть никогда не попробует ходы, которые она изначально (случайно, из-за начальной инициализации весов) оценила как плохие, и никогда не узнает, что на самом деле они хороши. Чтобы этого избежать, используется ε-greedy стратегия:
с вероятностью
εход выбирается случайно (исследование - сеть пробует то, что обычно не выбрала бы).с вероятностью (1 − ε) ход выбирается жадно, по текущей оценке сети (использование накопленных знаний).
Во всех рассмотренных реализациях ε линейно уменьшается по ходу обучения от высокого стартового значения (по-умолчанию 0.5) к низкому (по-умолчанию 0.02):
в начале обучения, когда сеть ещё ничего не знает об игре, высокий
εзаставляет сеть активно исследовать разнообразные позиции self-play, а не быстро зацикливаться на нескольких случайно "понравившихся" линиях игры.к концу обучения, когда оценки сети уже достаточно точны, низкий
εпозволяет сети в основном играть на пределе своих текущих знаний, лишь изредка отклоняясь для точечного дообследования.
Важно, что ε-greedy влияет и на то, какие данные видит алгоритм обучения, и (для on-policy методов) на характер самого target-а. Например, в Online TD(0) сеть оценивает ценность именно той политики, которая реально используется для игры, включая случайные исследовательские ходы (это как раз то, что отличает on-policy методы от off-policy, где Q-learning через max целится в оптимальную политику независимо от фактического ε).
2. Архитектура сети
Для всех методов обучения использована схема с одним скрытым слоем и tanh-для активации:
вход → полносвязный слой (tanh) → полносвязный слой (tanh) → выход
Сеть реализована на NumPy. Я не уверен, что это лучшая архитектура (простой полносвязный перцептрон) для подобной задачи, но в данном случае её вполне достаточно. Если есть интерес, то в программе довольно легко изменить структуру и состав слоёв.
Сигналы на выходе сети: −1 гарантированный проигрыш, 0 - ничья/нейтральная позиция, +1 гарантированный выигрыш.
Для V(s)-сети (Backward TD(0), Batch TD(0), Online TD(0), Monte Carlo): 9 входов, по одному на клетку доски. Клетка кодируется как +1 (X), −1 (O) или 0 (пусто). На скрытом слое 27 нейронов, на выходе один.
Для Q(s, a)-сети (Q-learning): 18 входов: 9 признаков состояния (доска, закодированная относительно текущего игрока, свои фишки всегда +1, чужие −1) плюс 9 признаков действия. На скрытом слое 36 нейронов, на выходе один.
3. Реализованные методы
Мне кажется, представленный ниже порядок методов наиболее удобен для изучения и понимания.
3.1 Backward Temporal Difference / Backward TD(0)
См. файл:
tic_tac_toe_backwardTD.py
Суть метода: Обновление весов отложено до конца партии. Сеть обучается по собственной оценке следующего состояния(bootstrapping) - обратным последовательным распространением TD-оценки.
Процесс обучения
Игра играется до конца, состояния записываются в
history = [S1,S2,...,Sn]Обновление идёт последовательно от
i = nкi = 1, причём каждое обновление сразу меняет веса сети:Для
i = n:target = result, веса обновляются.Для
i < n:target = V(S{i+1}). Значениеtargetберётся как текущая оценка сети следующего состояния (уже после обновления весов на более поздних ходах). Затем веса снова обновляются
Поясняющий пример
Допустим выиграл X. Пусть этому состоянию соответствует состояние S5, т.е.:
S1 → S2 → S3 → S4 → S5(победа). S5 соответствует result = 1.0.
Теперь эта оценка протягивается(бутстрепится) последовательно от конечного состояния к начальному:
S5 → target=result=1.0
S4 → target=V(S5)
S3 → target=V(S4)
S2 → target=V(S3)
S1 → target=V(S2)
При этом, для каждого промежуточного состояния выполняется коррекция TD-ошибки
текущее состояние Sᵢ │ ▼ V(Sᵢ) = net.forward(Sᵢ) # предсказание сети │ ▼ target = V(Sᵢ₊₁) # проброс оценки │ ▼ ошибка = V(Sᵢ) − target │ ▼ backprop (обновление весов сети)
3.2 Batch(episode) Temporal Difference / Batch TD(0)
См. файл:
tic_tac_toe_batchTD.py
Суть метода: Обновление весов отложено до конца партии. Сеть обучается по собственной оценке следующего состояния(bootstrapping), но распространение TD-оценки выполняется последовательно от начального состояния к терминальному(в отличие от предыдущего метода).
Процесс обучения
Игра играется до конца, состояния записываются в
history = [S1,S2,...,Sn].Обновление идёт последовательно от
i = 1кi = n, причём каждое обновление сразу меняет веса сети:Для
i = n(последнее состояние):target = result.Для
i < n:target = V(S{i+1}). Значениеtargetберётся как текущая оценка сети следующего состояния (на момент вычисленияtargetвеса уже могут быть обновлены предыдущими шагами этой же партии).
Отличие от backward-версии
В этом методе target для Si вычисляется до того, как сеть что-либо узнала о состоянии S{i+1} через обновление на этом шаге эпизода -V(S{i+1}) берется как есть, без переиспользования свежей информации, полученной чуть позже в этом же проходе. В частности:
targetдляS1вычисляется на весах, которые уже слегка изменены (от прошлых эпизодов обучения), но не отражают информацию об исходе именно этой партии, т.к. эта информация (т.е.result) будет вплетена в веса только на последнем шаге текущего прохода (т.е. наi = n), уже после того какS1обновлён.Из-за этого сигнал о реальном исходе партии распространяется от
SnкS1не за один проход, а постепенно, по одному шагу за эпизод, т.е. требуется много партий, чтобы это влияние диффузно дошло до самых ранних ходов.
3.3 Monte Carlo
См. файл:
tic_tac_toe_MC.py
Суть метода: Сеть не использует собственные промежуточные оценки позиций/состояний. Модель обучается исключительно на реальном исходе партии без bootstrapping, без предположений/оценок о промежуточных состояниях.
Процесс обучения
Игра играется до конца, состояния записываются в
history = [S1,S2,...,Sn].После завершения partии происходит обновление всех состояний одним и тем же target-ом:
для любого
Si(включая терминальноеSn):target = result
3.4 Online TD(0)
См. файл:
tic_tac_toe_TD.py
Суть метода: Классический TD(0) в "онлайн" форме, т.е. без ожидания окончания партии. Обновление весов выполняется после каждого хода - по одной обновляемой паре (Si, target) сразу по мере того, как становится известен S{i+1}.
Процесс обучения
Игра играется ход за ходом. После каждого хода запоминается
prev_state(состояние до текущего хода) иcurr_state(состояние после текущего хода, оно же состояние до следующего хода противника).Обучение происходит не по сохраненной истории после окончания партии, а онлайн, т.е. прямо по ходу игры - сразу после того, как становится известен очередной
curr_state:Если
curr_stateтерминальный (т.е. игра закончена), тоtarget = result, и на этом target-е обновляетсяprev_state.Если
curr_stateне терминальный, тоtarget = V(curr_state)(оценка сети для следующего состояния), и на этом target-е тоже обновляетсяprev_state.
Первое состояние партии (
S1, после первого хода X) не участвует как объект обновления сразу - оно становитсяprev_stateи обучается только на следующей итерации, когда появитсяS2. Аналогично последнее состояние партии (Sn, терминальное) никогда не подаётся вtrain_stepкак обучаемый вход, оно используется только как источникtarget = resultдля предпоследнего состоянияS{n-1}.
3.5 Q-Learning
См. файл:
tic_tac_toe_QL.py
Суть метода: Оценка функции ценности действия(action-value function) в состоянии Q(s, a), т.е. хода a в позиции S. Q(s, a) - ожидаемый итоговый результат партии, если в состоянии S сделать ход a, а дальше играть оптимально (с точки зрения текущего игрока). Обновление весов происходит сразу после каждого хода (онлайн). Сеть обучается с помощью bootstrapping, target строится на основе максимальной оценки следующего состояния с точки зрения противника.
Поскольку target строится через максимум по всем возможным ходам противника (а не по тому ходу, который противник реально сделает дальше с учётом собственного ε), метод является off-policy - сеть напрямую обучается оптимальной Q-функции, независимо от того, насколько шумной(случайной, исследовательской) была фактическая политика self-play, использованная для сбора партий.
Процесс обучения
Игра идёт ход за ходом. После каждого хода текущего игрока сразу выполняется обновление.
Формирование
targetи обновление:Если ход привёл к победе текущего игрока, то:
target = +1Если ничья:
target = 0Иначе:
target = −max_{a'} Q(s', a'), гдеs'- позиция после хода, закодированная уже с точки зрения противника, аmaxберётся по всем допустимым ходам противника изs'.
После вычисления target веса сети сразу обновляются.
4. Заключение
Все пять методов решают одну и ту же задачу, а именно - вывести из партий self-play функцию ценности. Конечно меоды решают задачу с разной скоростью, разной эффективностью. Может быть для других более сложных игр эти методы и вовсе не годятся. Тем не менее мы познакомились с базовыми подходами на которых строится современное обучения с подкреплением.
Исходный код методов: исходный код на github

