В статье я попробую осветить подход к обучению нейронной сети игре в крестики-нолики с помощью методов обучения с подкреплением (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
Автор: Dima_gangsta


