본문으로 건너뛰기

RL 008

· 약 9분

RL Taxonomy​

Drawbacks of previous methods​

  • Large state space:
    • Go: 1017010^{170}
    • Backgammon: 102010^{20}
    • Atari games: 109 to 101110^9 \text{ to } 10^{11}
  • Scale-up model-free based techniques for prediction and control is challenging.
  • Value function used look-up table representation
    • All states has an entry in V(s)V(s)
    • All state-action pair (s,a)(s,a) has an entry in Q(s,a)Q(s,a)
  • Too many state, and stat-action pair to store in memory
  • Slow learning process for each states, due to large space.
  • Each states needs to be explored sufficiently.
  • For large MDPs
    • Function approximation: Almost near optimal value.
    • Estimate value function with function approximation.
      • v^(s,w)≈vπ(s)\hat{v}(s,w) \approx v_\pi(s)
      • q^(s,a,w)≈qπ(s,a)\hat{q}(s,a,w) \approx q_\pi(s,a)
    • Generalize from seen states to unseen states.
    • Using MC or TD learning techniques: Update parameter ww.

Function Approximation​

  • A technique for estimating unknown underlying function using historical or available observations.
  • Assumption: an underlying mapping function exists
    • f(x)≈f^(x)f(x) \approx \hat{f}(x)
  • Function: x→fyx \xrightarrow{f} y

Types of Value Function Approximation​

State-value: state ss → scalar v^(s,w)\hat{v}(s,w)

Action-value (s, a input): state ss and action aa → scalar q^(s,a,w)\hat{q}(s,a,w)

Action-value (s input, all actions out): state ss → q^(s,a1,w),…,q^(s,am,w)\hat{q}(s,a_1,w),\ldots,\hat{q}(s,a_m,w)

  • ww is the parameter vector of the function approximator.
    • e.g. Neural Network weights

Types of Function Approximator​

  • Tabular
    • V-Table: s→v(s)s \rightarrow v(s)
    • Q-Table: (s,a)→q(s,a)(s,a) \rightarrow q(s,a)
  • Decision Trees, Nearest Neighbors
    • s=[speed=49,distance=10]s = [\text{speed} = 49, \text{distance} = 10]
    • s→Decision Tree / KNN→v^(s)s \rightarrow \text{Decision Tree / KNN} \rightarrow \hat{v}(s)
  • Linear Function approximation
    • Linear Combination of Features
    • Values are linear function for features v^(s,w)=wTx(s)\hat{v}(s,w) = w^Tx(s)
    • x(s)=[speeddistancefuel]x(s) = \begin{bmatrix} \text{speed} \\ \text{distance} \\ \text{fuel} \end{bmatrix}
    • v^(s,w)=wTx(s)\hat{v}(s,w) = w^Tx(s)
    • v^(s,w)=w1x1(s)+⋯+wnxn(s)\hat{v}(s,w) = w_1x_1(s) + \cdots + w_n x_n(s)
  • Differentiable function approximation
    • v^(s,w)\hat{v}(s,w) is a differentiable function of ww, can be non-linear in ss
    • e.g. Neural Network, or CNN
      • ∂v^(s,w)∂w\frac{\partial \hat{v}(s,w)}{\partial w}
      • w←w−α∇wLw \leftarrow w - \alpha \nabla_wL
      • Pixels→CNNv^(s)\text{Pixels} \xrightarrow{CNN} \hat{v}(s)

Which FA to use?​

In principle, any FA that fits the RL framework can be used.

FANotes
TabularEasy; not scalable; does not generalise
LinearRequires good features
Differentiable (better choice)Scalable; not always well understood
Neural NetworksPerforms well
Deep Neural NetworksPopular choice; performs well
  • Need training methods that are suitable for non-stationary data.

ANN​

Logistic Regression​

L(a,y)=−ylog⁡(a)−(1−y)log⁡(1−a)\mathcal{L}(a, y) = -y \log(a) - (1-y) \log(1-a)

  • Loss function for Logistic Regression

Gradient Descent​

Goal: Find ww tht minimize J(w)J(w) or find a local minimum of J(w)J(w).

  • An iterative approach for error correction.
  • For one sample, L(a,y)=−ylog⁡(a)−(1−y)log⁡(1−a)\mathcal{L}(a, y) = -y \log(a) - (1-y) \log(1-a)
  • For mm samples,

J(w,b)=1m∑i=1mL(a(i),y(i))J(w, b) = \frac{1}{m} \sum_{i=1}^m \mathcal{L}(a^{(i)}, y^{(i)})

  • Find ww and bb that will minimize J(w,b)J(w, b)
    • min⁡w,bJ(w,b)\min_{w, b} J(w, b)
    • J(w)J(w) is loss function.

w←w−α∂J(w,b)∂w,b←b−α∂J(w,b)∂bw \leftarrow w - \alpha \frac{\partial J(w, b)}{\partial w}, \quad b \leftarrow b - \alpha \frac{\partial J(w, b)}{\partial b}

ΔwJ(w)=(∂J(w)∂w1,…,∂J(w)∂wn)\Delta w J(w) = \big(\frac{\partial J(w)}{\partial w_1}, \ldots, \frac{\partial J(w)}{\partial w_n}\big)

  • Adjust the ww in the direction of the -ve gradient.
    • w=w−αΔw,b=b−αΔbw = w - \alpha \Delta w, \quad b = b - \alpha \Delta b
    • where α\alpha is the learning rate.
    • and, Δw=∂J(w,b)∂w,Δb=∂J(w,b)∂b\Delta w = \frac{\partial J(w, b)}{\partial w}, \quad \Delta b = \frac{\partial J(w, b)}{\partial b}

ANN as a Function Approximator​

  • Need a method to estimate: Δw\Delta w
  • Method to incrementally adjust and update: ww
  • Representation of some aspects of RL environments as feature vector x(s)x(s)
  • Aspects can be approximated: V(s) or Q(s,a)V(s) \text{ or } Q(s, a)
    • Estimate value functions with function approximation.
    • V^(s,w)≈V(s)\hat{V}(s, \color{red}{w}\color{black}) \approx V(s)
    • Q^(s,a,w)≈Q(s,a)\hat{Q}(s, a, \color{red}{w}\color{black}) \approx Q(s, a)
  • A mechanism to plugin MC and TD for function approximation.

Value Function Approximation for policy evaluation​

VFA: Value Function Approximation

  • vπ(s)v_\pi(s) is given, so making it supervised.
  • v^(s,w)\hat{v}(s, w) is estimated value if it follows the policy π\pi by using ANN.
  • x→y^vsyx \rightarrow \hat{y} \quad \text{vs} \quad y (supervised learning)
  • s→v^(s,w)vsvπ(s)s \rightarrow \hat{v}(s, w) \quad \text{vs} \quad v_\pi(s) (Reinforcement learning)

J(w)=Eπ[(vπ(s)−v^(s,w))2] J(w) = \mathbb{E}_\pi[\big(v_\pi(s) - \hat{v}(s, w)\big)^2]

  • Find ww that minimizes J(w)J(w)
    • J(w)J(w): Loss function.
    • v^(s,w)\hat{v}(s, w): Guess from the value function approximation.
    • vπ(s)−v^(s,w)v_\pi(s) - \hat{v}(s, w): Error between the true value and the estimated value.
    • Eπ[⋅]\mathbb{E}_\pi[\cdot]: Expected average over all states ss according to the policy π\pi.
Δw=−12α∇wJ(w)Δw=αEπ[(vπ(s)−v^(s,w))⏟Error∇wv^(s,w)⏟Gradient]\begin{aligned} & \Delta w = -\frac{1}{2} \alpha \nabla_w J(w) \\ & \Delta w = \alpha \mathbb{E}_\pi[\underbrace{\big(v_\pi(s) - \hat{v}(s, w)\big)}_{Error} \underbrace{\nabla_w \hat{v}(s, w)}_{Gradient}] \end{aligned}
  • Δw=\Delta w = Step size X Error X Gradient

VFA with SGD​

Δw=α(vπ(s)−v^(s,w))∇wv^(s,w) \Delta w = \alpha (v_\pi(s) - \hat{v}(s, w)) \nabla_w \hat{v}(s, w)

  • Stochastic Gradient Descent sample the gradient
  • Sample a state randomly, find what vπ(s)v_\pi(s) is, and find the estimate
    • Δw=α[Error at s×Gradient at s]\Delta w = \alpha [\text{Error at } s \times \text{Gradient at } s]
    • Sample sts_t from π\pi, and update ww
  • Expected value update using SGD update is same as full gradient update
    • Random sampling will be eventually the same as full gradient update

Linear Value Function Approximation​

  • Represent Value function (state or action) by a linear combination of features
    • v^(s,w)=x(s)Tw=∑j=0nxj(s)wj\hat{v}(s, w) = x(s)^Tw = \sum_{j=0}^n x_j(s) w_j
  • Loss function: J(w)=Eπ[(vπ(s)−v^(s,w))2]J(w) = \mathbb{E}_\pi[\big(v_\pi(s) - \hat{v}(s, w)\big)^2]

Δw=α(vπ(s)−v^(s,w))⏟Errorx(s)⏟Feature\Delta w = \alpha \underbrace{\big(v_\pi(s) - \hat{v}(s, w)\big)}_{Error} \underbrace{x(s)}_{Feature}

  • Weight update: step size X Error X Feature value
  • SGD converges to global mininum

Incremental methods for prediction​

  • Real-world RL is not supervised, instead, it is guided by rewards.
  • No vπ(s)v_\pi(s) is given.
  • Instead, substitute a target for vπ(s)v_\pi(s)

Δw=α(Gt−v^(st,w))∇wv^(st,w) \Delta w = \alpha(G_t - \hat{v}(s_t, w)) \nabla_w \hat{v}(s_t, w)

  • Monte Carlo updates: target is the return GtG_t

Δw=α(Rt+1+γv^(st+1,w)⏟TD Target−v^(st,w))∇wv^(st,w) \Delta w = \alpha(\underbrace{R_{t+1} + \gamma \hat{v}(s_{t+1}, w)}_{\text{TD Target}} - \hat{v}(s_t, w)) \nabla_w \hat{v}(s_t, w)

  • TD updates: use TD Target

VFA for Monte Carlo​

  • MC return GtG_t is an unbiased, noisy sample of true value vπ(st)v_\pi(s_t)
    • if Gt=3,7,4,6,5,⋯G_t = 3, 7, 4, 6, 5, \cdots
    • but in average, 1N∑i=1NGt(i)≈vπ(s)\frac{1}{N} \sum_{i=1}^N G_t^{(i)} \approx v_\pi(s)
  • ⟨S1,G1⟩,⟨S2,G2⟩,⋯ ,⟨SN,GN⟩\langle S_1, G_1 \rangle, \langle S_2, G_2 \rangle, \cdots, \langle S_N, G_N \rangle is an estimate of vπ(s)v_\pi(s)
    • x→yx \rightarrow y (supervised learning)
    • st→Gts_t \rightarrow G_t (Reinforcement learning)
  • so it learns to v^(s,w)≈Gt\hat{v}(s, w) \approx G_t
Δw=α(Gt−v^(s,w))∇wv^(s,w)Δw=α(Gt−v^(s,w))x(s)\begin{aligned} & \Delta w = \alpha(G_t - \hat{v}(s, w)) \nabla_w \hat{v}(s, w) \\ & \Delta w = \alpha \big(G_t - \hat{v}(s, w)\big) x(s) \end{aligned}
  • In Linear VFA, v^(s,w)=wTx(s)\hat{v}(s,w) = w^Tx(s)
    • x(s)x(s) is the feature vector of state ss
    • ww is the weight vector
    • so ∇wv^(s,w)=∇w(wTx(s))=x(s)\nabla_w \hat{v}(s, w) = \nabla_w (w^Tx(s)) = x(s)
  • Update = Learning rate X Error X state features
Initialize w=0,k=1Loop:Sample k-th episode (sk1,ak1,rk1,...,SkLk) using policy πFor t=1,...Lk:if First Visit to (s) in episode k,thenGt(s)=∑j=tLkγk,jWeight update: w←w−α(Gt−v^(s,w))x(s)k=k+1\begin{aligned} & \text{Initialize } w = 0, k = 1 \\ & Loop: \\ & \quad \text{Sample k-th episode } (s_{k1}, a_{k1}, r_{k1}, ..., S_{kL_k}) \text{ using policy } \pi \\ & \quad \text{For } t = 1, ... L_k: \\ & \quad \quad \text{if First Visit to } (s) \text{ in episode } k, then \\ & \quad \quad \quad G_t(s) = \sum_{j=t}^{L_k} \gamma_{k,j} \\ & \quad \quad \quad \text{Weight update: } \quad w \leftarrow w - \alpha(G_t - \hat{v}(s,w))x(s) \\ & \quad \quad k = k + 1 \\ & \end{aligned}
  • Gt=Rt+1+γRt+2+γ2Rt+3+⋯+γLk−tRLkG_t = R_{t+1} + \gamma R_{t+2} + \gamma^2 R_{t+3} + \cdots + \gamma^{L_k-t} R_{L_k}

VFA for TD Learning​

  • Use bootstrapping and sampling to approximate the target vπv_\pi
    • Sampling: Use a sampled transition St,At,Rt+1,St+1S_t, A_t, R_{t+1}, S_{t+1} instead of computing the full expectation over possible transitions.
    • Bootstrapping: Construct the TD target using the current estimate v^(St+1,w)\hat{v}(S_{t+1}, w), instead of waiting for the complete episodic return.

Rt+1+γv^(St+1,w)R_{t+1} + \gamma \hat{v}(S_{t+1}, w)

  • TD Target is a biased sample of true value vπ(St)v_\pi(S_t)
  • Can use supervised learning on a training set of ⟨S,R⟩\langle S, R \rangle paris
    • ⟨S1,R2+γv^(S2,w)⟩,⟨S2,R3+γv^(S3,w)⟩,⋯ ,⟨ST−1,RT⟩\langle S_1, R_2 + \gamma \hat{v}(S_2, w) \rangle, \langle S_2, R_3 + \gamma \hat{v}(S_3, w) \rangle, \cdots, \langle S_{T-1}, R_T \rangle

Δw=α(R+γv^(s′,w)−v^(s,w))∇wv^(s,w) \Delta w = \alpha(R + \gamma \hat{v}(s', w) - \hat{v}(s, w)) \nabla_w \hat{v}(s, w)

  • Linear VFA + TD(0) for prediction and policy evaluation
  • TD Error: δt=Rt+1+γv^(St+1,w)−v^(St,w)\delta_t = R_{t+1} + \gamma \hat{v}(S_{t+1}, w) - \hat{v}(S_t, w)
  • Δw=αδt∇wv^(St,w)\Delta w = \alpha \delta_t \nabla_w \hat{v}(S_t, w)
    • ∇wv^(s,w)=x(s)\nabla_w \hat{v}(s, w) = x(s)
    • Δw=αδtx(s)\Delta w = \alpha \delta_t x(s)
    • w←w−αδtx(s)w \leftarrow w - \alpha \delta_t x(s)
Initialize w=0Loop:Initialize sSample transition (s,a,r,s′) using policy πWeight update: w←w−α(R+γv^(s′,w)−v^(s,w))x(s)\begin{aligned} & \text{Initialize } w = 0 \\ & Loop: \\ & \quad \text{Initialize } s \\ & \quad \text{Sample transition } (s, a, r, s') \text{ using policy } \pi \\ & \quad \text{Weight update: } \quad w \leftarrow w - \alpha(R + \gamma \hat{v}(s', w) - \hat{v}(s, w))x(s) \\ & \end{aligned}

Overall​

  • Linear VFA: v^(s,w)=wTx(s)\hat{v}(s, w) = w^Tx(s)
  • MC Target = GtG_t
  • TD Target = Rt+1+γv^(s′,w)R_{t+1} + \gamma \hat{v}(s', w)
  • MC: w←w+α(Gt−v^(s,w))x(s)w \leftarrow w + \alpha(G_t - \hat{v}(s, w)) x(s)
  • TD: w←w+α(Rt+1+γv^(s′,w)−v^(s,w))x(s)w \leftarrow w + \alpha(R_{t+1} + \gamma \hat{v}(s', w) - \hat{v}(s, w))x(s)