Trust Region Policy Optimization

TRPO 논문리뷰

Abstract

Introduction

Background

flowchart TD
	subgraph GB1[Policy Iteration]
	GB1A[ADP]
	GB1B[CPI]
	end
	
	subgraph GB2[Policy Gradient]
	direction TB
	GB2A[REINFORCE]
	GB2B[NPG]
	GB2C[Actor-Critic]
	GB2A ~~~ GB2B
	GB2A ~~~ GB2C
	end
	
	subgraph DFO[Derivative-Free]
	DFO1[CEM]
	DFO2[CMA]
	end
	
	GB[Gradient-Based]
	GB --- GB1
	GB --- GB2
	
	A["Policy Optimization<br/>Algorithm"]
	A === GB
	A === DFO

Preliminary

Notations

Symbol Definition Name 의미
\(\mathcal{S}\) \(\mathcal{S} = \{s: \text{state}\}\) Finite set of states 유한상태공간
\(\mathcal{A}\) \(\mathcal{A} = \{a: \text{action}\}\) Finite set of actions 유한행동공간
\(P(s' \mid s, a)\) \(P: \mathcal{S} \times \mathcal{A} \times \mathcal{S} \to \mathbb{R}\) Transition probability distribution 상태 \(s\) 에서 행동 \(a\) 후 상태 \(s'\) 로 전이 확률
\(r(s)\) \(r: \mathcal{S} \to \mathbb{R}\) Reward function 상태 \(s\) 의 보상
\(\rho_{0}(s)\) \(\rho_0: \mathcal{S} \to \mathbb{R}\) The distribution of the initial state 초기상태 \(s_0\) 의 분포
\(\gamma\) \(\gamma \in (0, 1)\) Discount factor 미래 보상 할인 인자
\(\pi(a \mid s)\) \(\pi: \mathcal{S} \times \mathcal{A} \to [0,1]\) Stochastic policy 상태 \(s\) 에서 행동 \(a\) 를 선택할 확률
\(\tilde{\cdot}\) target, sample, revised, \(\cdots\) - 문맥에 따름
\(\eta(\pi)\) \(\mathbb{E}_{s_0, a_0, \cdots}\left[\sum_{t=0}^{\infty} \gamma^t r(s_t)\right]\) Expected discounted reward of stochastic policy \(\pi\) 정책 \(\pi\) 의 기대 할인 보상
\(Q_\pi(s_t,a_t)\) \(\mathbb{E}_{s_{t+1}, a_{t+1}, \cdots}\left[\sum_{l=0}^{\infty} \gamma^l r(s_{t+l})\right]\) State-action value function 상태 \(s_t\) 에서 특정 행동 \(a_t\) 에 대한 기대 보상
\(V_\pi(s_t)\) \(\mathbb{E}_{a_t, s_{t+1}, \cdots}\left[\sum_{l=0}^{\infty} \gamma^l r(s_{t+l})\right]\) \(=\) \(\sum_{a_t} \pi(a_t \mid s_t) Q_{\pi} (s_t, a_t)\) Value function 상태 \(s_t\) 에서 모든 행동에 대한 기대 보상 가중 평균
\(A_\pi(s,a)\) \(Q_\pi(s,a) - V_\pi(s)\) Advantage function 행동 \(a\) 가 평균 대비 얼마나 좋은가
\(\rho_\pi(s)\) \(\sum_{t=0}^{\infty} \gamma^t P(s_t = s)\) Discounted visitation frequency 할인이 적용된 상태 방문 빈도
\(\mathbb{D}_{TV}(p \;|\; q)\) \(\frac{1}{2} \sum_i \lvert{p_i - q_i}\rvert\) Total variation divergence 사건 확률 거리. 아예 안 겹치면 1
\(\mathbb{D}^{max }_{TV}(\pi, \tilde{\pi})\) \(\underset{s}{max}\; \mathbb{D}_{TV}(\pi(\cdot \mid s) \;|\; \tilde{\pi}(\cdot \mid s))\) Max total variation 모든 상태 중 최악(max)의 TV divergence
\(\mathbb{D}_{KL}(p \;|\; q)\) \(\sum_i p_i \log{\frac{p_i}{q_i}}\) KL divergence 분포 \(q\) 로 \(p\) 를 인코딩할 때 정보 손실
\(\mathbb{D}^{max }_{KL}(\pi, \tilde{\pi})\) \(\underset{s}{max}\; \mathbb{D}_{KL}(\pi(\cdot \mid s) \;|\; \tilde{\pi}(\cdot \mid s))\) Max KL divergence 모든 상태 중 최악(max)의 KL divergence
\(\bar{\mathbb{D}}^{\rho}_{KL}(\pi, \tilde{\pi})\) \(\mathbb{E}_{s \sim \rho}\left[\mathbb{D}_{KL}\bigl(\pi(\cdot \mid s) \;|\; \tilde{\pi}(\cdot\mid s)\bigr)\right]\) Average KL divergence 상태 분포 \(\rho\) 하의 평균 KL divergence
\(F_{ij}\) \(\mathbb{E}_{s \sim \rho_{\theta_{odd}}} \left[{\frac{{\partial}^2}{\partial\theta_i\partial\theta_j}} \mathbb{D}_{KL}\left(\pi_{\theta_{old}}(\cdot \mid s) \;|\; \pi_\theta(\cdot \mid s)\right)\right]\) Fisher information matrix (Hessian of KL) 해석적 FIM

Kullback-Leibler Divergence

Conservative Policy Iteration(CPI)

Natural Policy Gradient (NPG)

Overview

Surrogate objective function

\(L_\pi (\tilde{\pi}) = \eta (\pi) + \sum_{s} {\rho_{\pi} (s) \sum_{a} {\tilde{\pi} (a \vert s) A^{\pi}(s, a)}}\)

Expected return(\(\eta\)) of policy \(\tilde{\pi}\)

\(\eta(\tilde{\pi}) = \eta(\pi) + \sum_{s} {\rho_{\tilde{\pi}} (s) \sum_{a} {\tilde{\pi} (a \vert s) A^{\pi}(s, a)}}\)

The Theoretically Justified Algorithm

Minorization-Maximization (MM) algorithm

Approximations

1. max KL \(\approx\) avg KL

\(\mathbb{D}^{max}_{KL} (\theta_{old}, \theta) \approx \bar{\mathbb{D}}^{\rho_{\theta_{old}}}_{KL}\) \(\max_s D_{KL}\bigl(\pi(\cdot\vert s) \| \tilde{\pi}(\cdot\vert s)\bigr) \approx \mathbb{E}_{s \sim \rho}\left[D_{KL}\bigl(\pi(\cdot\vert s) \| \tilde{\pi}(\cdot\vert s)\bigr)\right]\)

2. KL penalty \(\to\) KL constraint

\(\underset{\theta}{\max}\left[L_{\theta_{old}}(\theta) - C \cdot D^{max}_{KL}(\theta_{old}, \theta)\right] \;\;\longrightarrow\;\; \underset{\theta}{\max}\; L_{\theta_{old}}(\theta) \;\;\text{subject to}\;\; \bar{D}^{\rho_{\theta_{old}}}_{KL}(\theta_{old}, \theta) \leq \delta\)

3. \(A_{\pi}\) estimation via Monte-Carlo sampling
Single-path

\(\hat{Q}(s_t, a_t) = \sum_{l=0}^{T-t} {\gamma}^{l} r(s_{t+l})\)

Vine

\(\hat{Q}(s_n, a_{n, k}) = r(s_n) + \gamma r(s_1') + {\gamma}^2 r(s_2') + \cdots\)

Trust Region Policy Optimization

Full Derivation

Objective

Surrogate

Monotonic Improvement Guarantee for General Stochastic Policies

Definition (\(\alpha\)-coupling).

A policy pair \((\pi, \tilde{\pi})\) is \(\alpha\)-coupled if there exists a joint distribution \((a, \tilde{a}) \mid s\) with marginals \(\pi(\cdot \mid s)\), \(\tilde{\pi}(\cdot \mid s)\) such that \(P(a \ne \tilde{a} \mid s) \leq \alpha\)

Proposition.

(Levin et al.) If \(\mathbb{D}_{TV}^{max} (\pi, \tilde{\pi}) \leq \alpha\), then an \(\alpha\)-coupling exists.

Theorem 1. Policy Improvement Bound

Let \(\alpha = \mathbb{D}_{TV}^{max} (\pi_{old}, \pi_{new})\). Then the following bound holds: \(\eta({\pi_{new}}) \geq L_{\pi_{old}} (\pi_{new}) - \frac{4\epsilon\gamma}{(1-\gamma)^2} {\alpha}^2\) \(where \;\epsilon = \underset{s, a}{max} \lvert{A_\pi (s, a)}\rvert\) Proof. Taking expectation over trajectories \(\tau := (s_0, a_0, s_1, a_1, \cdots)\), \(\eta(\tilde{\pi}) = \eta(\pi) + \mathbb{E}_{\tau \sim \tilde{\pi}} \left[{\sum_{t=0}^{\infty} {\gamma}^t A_{\pi} (s_t, a_t)}\right]\) Define the expected advantage of \(\tilde{\pi}\) over \(\pi\) at state \(s\): \(\bar{A}(s) = \mathbb{E}_{a \sim \tilde{\pi}(\cdot \mid s)} \left[{A_{\pi}(s, a)}\right]\) Note that \(L_{\pi}\) can be written as \(L_{\pi}(\tilde{\pi}) = \eta(\pi) + \mathbb{E}_{\tau \sim \pi} \left[{\sum_{t=0}^{\infty} {\gamma}^t \bar{A} (s_t)}\right]\) The difference between \(\eta(\tilde{\pi})\) and \(L_\pi\) : \(\eta(\tilde{\pi}) - L_{\pi} (\tilde{\pi}) = \mathbb{E}_{\tau \sim \tilde{\pi}} \left[{\sum_{t=0}^{\infty} {\gamma}^t A_{\pi} (s_t, a_t)}\right] - \mathbb{E}_{\tau \sim \pi} \left[{\sum_{t=0}^{\infty} {\gamma}^t \bar{A} (s_t)}\right] \tag{*}\) To ensure monotonic improvement, the absolute value of \((*)\) should be bounded.

Lemma 1. \(\lvert \bar{A}(s) \rvert \leq 2\alpha\epsilon\)

\(\bar{A}(s) = \mathbb{E}_{\tilde{\pi}} \left[A\pi(s, a)\right] - \mathbb{E}_{\pi} \left[A\pi(s, a)\right]\) (TODO)

Lemma 2. \(\mathbb{D}_{TV} (P^{\tilde{\pi}}_t, P^{\pi}_t) \leq t\alpha\)

Construct paired process \((s_t, \tilde{s}_t)\) sharing transition randomness under \(\alpha\)-coupling (TODO)

Bound timestep \(t\) term of \((*)\) \(\begin{aligned} \biggl| \mathbb{E}P^{\tilde{\pi}}_t \left[\bar{A}(s)\right] - \mathbb{E}P^{\pi}_t \left[\bar{A}(s)\right] \biggr| &= \biggl| \sum_s \left({P^{\tilde{\pi}}_t(s) - P^{\pi}_t}(s)\right) \bar{A}(s) \biggr| \\ &\leq \underset{s}{max} \left|\bar{A}(s)\right| \;\cdot\; \sum_s \left|{P^{\tilde{\pi}}_t(s) - P^{\pi}_t}(s)\right| \\ &= \underset{s}{max} \left|\bar{A}(s)\right| \;\cdot\; 2 \mathbb{D}_{TV} (P^{\tilde{\pi}}_t, P^{\pi}_t) \\ &\leq (2\alpha\epsilon)\;\cdot\;(2t\alpha) = 4\epsilon\alpha^2t \end{aligned}\) Total timesteps: \(\begin{aligned} \lvert \eta(\tilde{\pi}) - L_{\pi} (\tilde{\pi}) \rvert &= \sum_{t=0}^{\infty} \biggl| \mathbb{E}_{s \sim P^{\tilde{\pi}}_t} \left[\bar{A}(s)\right] - \mathbb{E}_{s \sim P^{\pi}_t} \left[\bar{A}(s)\right] \biggr| \\ &\leq \sum_{t=0}^{\infty} \gamma^t \cdot 4\epsilon\alpha^2t = 4\epsilon\alpha^2 \cdot \frac{\gamma}{(1-\gamma)^2} \end{aligned}\) Therefore, \(\eta(\tilde{\pi}) \geq L_{\pi} (\tilde{\pi}) - \frac{4\epsilon\gamma}{(1-\gamma)^2} \alpha^2 \tag{**}\)

Minorization-Maximization

Definition (Minorizer). A function \(M_i(\pi_i)\) called minorizer such that \(M_i(\pi_i) = \eta(\pi_i)\) \(M_i(\pi) \leq \eta(\pi) \quad \text{for } \forall\pi\) Rearranging \((**)\), \(\begin{aligned} \eta(\tilde{\pi}) \geq \underbrace{L_{\pi} (\tilde{\pi}) - C\cdot\mathbb{D}_{KL}^{max}(\pi, \tilde{\pi})}_{M_i(\tilde{\pi})} \end{aligned}\) MM Algorithm

  1. construct a minorizer \(M_i(\pi)\) under the current policy \(\pi_i\).
  2. \[\pi_{i+1} = \underset{\pi}{argmax} M_i(\pi)\]
  3. repeat. \(\eta(\pi_{i+1}) \geq M_i(\pi_{i+1}) \geq M_i(\pi_i) = \eta(\pi_i)\)

Optimization of Parameterized Policies

KL Divergence의 2차 근사

Objective의 1차 근사

\(L_{\theta_{old}}(\theta) \approx L_{\theta_{old}}(\theta_{old}) + g^T (\theta - \theta_{old})\) \(g = \nabla_\theta L_{\theta_{old}}(\theta) \big|_{\theta = \theta_{old}}\)

제약 최적화의 해: Natural Gradient

Conjugate Gradient (CG) 로 \(F^{-1}g\) 근사

Fisher-Vector Product

\(Fv = \nabla_\theta \left[(\nabla_\theta \mathbb{D}_{KL})^T v\right]\)

        def fvp(params, v, states):

            def kl_fn(p):

                return avg_kl(p, params_old, states)

            g = jax.grad(kl_fn)(params)

            # g^T v 를 다시 미분

            return jax.grad(lambda p: jnp.dot(jax.grad(kl_fn)(p), v))(params)

        ```         또는 더 효율적으로 jax.jvp + jax.vjp 조합 사용.

CG Algorithm

Sample-Based Estimation of the Objective and Constraint

Importance Sampling

여기서 \(q = \pi_{\theta_{old}}\).

\[\hat{L}_{\theta_{old}}(\theta) = \frac{1}{|D|} \sum_{(s, a) \in D} \frac{\pi_\theta(a \mid s)}{\pi_{\theta_{old}}(a \mid s)} \hat{A}_{\theta_{old}}(s, a)\] \[\hat{\bar{\mathbb{D}}}_{KL} = \frac{1}{|D_s|} \sum_{s \in D_s} \mathbb{D}_{KL}\left(\pi_{\theta_{old}}(\cdot \mid s) \;\|\; \pi_\theta(\cdot \mid s)\right)\]

    - KL divergence는 closed-form (Gaussian policy의 경우):

    $$\mathbb{D}_{KL}(\mathcal{N}_1 | \mathcal{N}_2) = \frac{1}{2}\left[\log\frac{ \Sigma_2 }{ \Sigma_1 } - d + \text{tr}(\Sigma_2^{-1}\Sigma_1) + (\mu_2-\mu_1)^T \Sigma_2^{-1} (\mu_2-\mu_1)\right]$$

    - 대각 공분산 \(\Sigma = \text{diag}(\sigma_1^2, \ldots, \sigma_d^2)\) 이면 더 단순해짐.

Advantage 추정: Generalized Advantage Estimation (GAE)

TD Residual

\(\delta_t^V = r(s_t) + \gamma V(s_{t+1}) - V(s_t)\)

Policy Gradient의 샘플 추정

Fisher-Vector Product의 샘플 추정

\(\hat{F}v = \frac{1}{|D_s|} \sum_{s \in D_s} \nabla_\theta\left[(\nabla_\theta \mathbb{D}_{KL})^T v\right]\)

Practical Algorithm

TRPO Full Procedure

for iteration = 1, 2, ... do
    1. Rn policy π_θ_old to collect trajectories D = {τ_1, ..., τ_N}
    2. Estimate avantages Â_t using GAE(γ, λ) with fitted value function V_φ
    3. Compute policy gradient: g = ∇_θ L̂(θ)|_{θ=θ_old}
    4. Use CG to compute: s ≈ F⁻¹g  (K ≈ 10 iterations)
    5. Compute max step: β = √(2δ / sᵀFs)
    6. Backtracking line search:
       θ_new = θ_old + c^j · β · s
       where j is the smallest integer such that:
         - KL(θ_old, θ_new) ≤ δ
         - L̂(θ_new) ≥ L̂(θ_old)
    7.Update value function V_φ by regression on collected returns
    8. θ_old ← θ_new
end for

주요 하이퍼파라미터

하이퍼파라미터 의미 전형적 값
\(\delta\) Trust region 크기 (KL constraint) \(0.01\)
\(\gamma\) Discount factor \(0.99\)
\(\lambda\) GAE parameter \(0.97\)
\(K\) CG iterations \(10\)
damping FIM regularization \(\lambda_{\text{damp}}\) \(0.1\)
backtrack coeff Line search 축소 비율 \(0.5\)
backtrack iters Line search 최대 반복 \(10\)

TRPO vs. Natural Policy Gradient (NPG)

| | NPG | TRPO | | ————- | ————————- | —————————————– | | 업데이트 | \(\theta + \alpha F^{-1}g\) | \(\theta + \beta F^{-1}g\) with line search | | Step size | 고정 \(\alpha\) | Adaptive (trust region \(\delta\) 기반) | | KL constraint | 없음 (implicit) | 명시적 (\(\bar{\mathbb{D}}_{KL} \leq \delta\)) | | 안정성 | \(\alpha\) 에 민감 | Line search로 보정, 더 robust |