Chapter 03
Policy Iteration
NotebookPython 37 cells
In [5]python · cell 1
python
import numpy as np
import pprint
import sys
if "../" not in sys.path:
sys.path.append("../")
from lib.envs.gridworld import GridworldEnvIn [6]python · cell 2
python
pp = pprint.PrettyPrinter(indent=2)
env = GridworldEnv()In [7]python · cell 3
python
# Taken from Policy Evaluation Exercise!
def policy_eval(policy, env, discount_factor=1.0, theta=0.00001):
"""
Evaluate a policy given an environment and a full description of the environment's dynamics.
Args:
policy: [S, A] shaped matrix representing the policy.
env: OpenAI env. env.P represents the transition probabilities of the environment.
env.P[s][a] is a list of transition tuples (prob, next_state, reward, done).
env.nS is a number of states in the environment.
env.nA is a number of actions in the environment.
theta: We stop evaluation once our value function change is less than theta for all states.
discount_factor: Gamma discount factor.
Returns:
Vector of length env.nS representing the value function.
"""
# Start with a random (all 0) value function
V = np.zeros(env.nS)
while True:
delta = 0
# For each state, perform a "full backup"
for s in range(env.nS):
v = 0
# Look at the possible next actions
for a, action_prob in enumerate(policy[s]):
# For each action, look at the possible next states...
for prob, next_state, reward, done in env.P[s][a]:
# Calculate the expected value
v += action_prob * prob * (reward + discount_factor * V[next_state])
# How much our value function changed (across any states)
delta = max(delta, np.abs(v - V[s]))
V[s] = v
# Stop evaluating once our value function change is below a threshold
if delta < theta:
break
return np.array(V)In [13]python · cell 4
python
def policy_improvement(env, policy_eval_fn=policy_eval, discount_factor=1.0):
"""
Policy Improvement Algorithm. Iteratively evaluates and improves a policy
until an optimal policy is found.
Args:
env: The OpenAI envrionment.
policy_eval_fn: Policy Evaluation function that takes 3 arguments:
policy, env, discount_factor.
discount_factor: gamma discount factor.
Returns:
A tuple (policy, V).
policy is the optimal policy, a matrix of shape [S, A] where each state s
contains a valid probability distribution over actions.
V is the value function for the optimal policy.
"""
# Start with a random policy
policy = np.ones([env.nS, env.nA]) / env.nA
while True:
# Implement this!
break
return policy, np.zeros(env.nS)In [14]python · cell 5
python
policy, v = policy_improvement(env)
print("Policy Probability Distribution:")
print(policy)
print("")
print("Reshaped Grid Policy (0=up, 1=right, 2=down, 3=left):")
print(np.reshape(np.argmax(policy, axis=1), env.shape))
print("")
print("Value Function:")
print(v)
print("")
print("Reshaped Grid Value Function:")
print(v.reshape(env.shape))
print("")
Output
Policy Probability Distribution: [[ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25] [ 0.25 0.25 0.25 0.25]] Reshaped Grid Policy (0=up, 1=right, 2=down, 3=left): [[0 0 0 0] [0 0 0 0] [0 0 0 0] [0 0 0 0]] Value Function: [ 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0.] Reshaped Grid Value Function: [[ 0. 0. 0. 0.] [ 0. 0. 0. 0.] [ 0. 0. 0. 0.] [ 0. 0. 0. 0.]]
In [15]python · cell 6
python
# Test the value function
expected_v = np.array([ 0, -1, -2, -3, -1, -2, -3, -2, -2, -3, -2, -1, -3, -2, -1, 0])
np.testing.assert_array_almost_equal(v, expected_v, decimal=2)Output
[0;31m---------------------------------------------------------------------------[0m
[0;31mAssertionError[0m Traceback (most recent call last)
[0;32m<ipython-input-15-55581f8eb5c9>[0m in [0;36m<module>[0;34m()[0m
[1;32m 1[0m [0;31m# Test the value function[0m[0;34m[0m[0;34m[0m[0m
[1;32m 2[0m [0mexpected_v[0m [0;34m=[0m [0mnp[0m[0;34m.[0m[0marray[0m[0;34m([0m[0;34m[[0m [0;36m0[0m[0;34m,[0m [0;34m-[0m[0;36m1[0m[0;34m,[0m [0;34m-[0m[0;36m2[0m[0;34m,[0m [0;34m-[0m[0;36m3[0m[0;34m,[0m [0;34m-[0m[0;36m1[0m[0;34m,[0m [0;34m-[0m[0;36m2[0m[0;34m,[0m [0;34m-[0m[0;36m3[0m[0;34m,[0m [0;34m-[0m[0;36m2[0m[0;34m,[0m [0;34m-[0m[0;36m2[0m[0;34m,[0m [0;34m-[0m[0;36m3[0m[0;34m,[0m [0;34m-[0m[0;36m2[0m[0;34m,[0m [0;34m-[0m[0;36m1[0m[0;34m,[0m [0;34m-[0m[0;36m3[0m[0;34m,[0m [0;34m-[0m[0;36m2[0m[0;34m,[0m [0;34m-[0m[0;36m1[0m[0;34m,[0m [0;36m0[0m[0;34m][0m[0;34m)[0m[0;34m[0m[0m
[0;32m----> 3[0;31m [0mnp[0m[0;34m.[0m[0mtesting[0m[0;34m.[0m[0massert_array_almost_equal[0m[0;34m([0m[0mv[0m[0;34m,[0m [0mexpected_v[0m[0;34m,[0m [0mdecimal[0m[0;34m=[0m[0;36m2[0m[0;34m)[0m[0;34m[0m[0m
[0m
[0;32m/Users/dennybritz/venvs/tf/lib/python3.5/site-packages/numpy/testing/utils.py[0m in [0;36massert_array_almost_equal[0;34m(x, y, decimal, err_msg, verbose)[0m
[1;32m 914[0m assert_array_compare(compare, x, y, err_msg=err_msg, verbose=verbose,
[1;32m 915[0m [0mheader[0m[0;34m=[0m[0;34m([0m[0;34m'Arrays are not almost equal to %d decimals'[0m [0;34m%[0m [0mdecimal[0m[0;34m)[0m[0;34m,[0m[0;34m[0m[0m
[0;32m--> 916[0;31m precision=decimal)
[0m[1;32m 917[0m [0;34m[0m[0m
[1;32m 918[0m [0;34m[0m[0m
[0;32m/Users/dennybritz/venvs/tf/lib/python3.5/site-packages/numpy/testing/utils.py[0m in [0;36massert_array_compare[0;34m(comparison, x, y, err_msg, verbose, header, precision)[0m
[1;32m 735[0m names=('x', 'y'), precision=precision)
[1;32m 736[0m [0;32mif[0m [0;32mnot[0m [0mcond[0m[0;34m:[0m[0;34m[0m[0m
[0;32m--> 737[0;31m [0;32mraise[0m [0mAssertionError[0m[0;34m([0m[0mmsg[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 738[0m [0;32mexcept[0m [0mValueError[0m[0;34m:[0m[0;34m[0m[0m
[1;32m 739[0m [0;32mimport[0m [0mtraceback[0m[0;34m[0m[0m
[0;31mAssertionError[0m:
Arrays are not almost equal to 2 decimals
(mismatch 87.5%)
x: array([ 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.,
0., 0., 0.])
y: array([ 0, -1, -2, -3, -1, -2, -3, -2, -2, -3, -2, -1, -3, -2, -1, 0])In [ ]python · cell 7
python
