Chapter 02
Policy Evaluation
NotebookPython 36 cells
In [23]python · cell 1
python
import numpy as np
import sys
if "../" not in sys.path:
sys.path.append("../")
from lib.envs.gridworld import GridworldEnvIn [24]python · cell 2
python
env = GridworldEnv()In [25]python · cell 3
python
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:
# TODO: Implement!
break
return np.array(V)In [26]python · cell 4
python
random_policy = np.ones([env.nS, env.nA]) / env.nA
v = policy_eval(random_policy, env)In [22]python · cell 5
python
# Test: Make sure the evaluated policy is what we expected
expected_v = np.array([0, -14, -20, -22, -14, -18, -20, -20, -20, -20, -18, -14, -22, -20, -14, 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-22-235f39fb115c>[0m in [0;36m<module>[0;34m()[0m
[1;32m 1[0m [0;31m# Test: Make sure the evaluated policy is what we expected[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;36m14[0m[0;34m,[0m [0;34m-[0m[0;36m20[0m[0;34m,[0m [0;34m-[0m[0;36m22[0m[0;34m,[0m [0;34m-[0m[0;36m14[0m[0;34m,[0m [0;34m-[0m[0;36m18[0m[0;34m,[0m [0;34m-[0m[0;36m20[0m[0;34m,[0m [0;34m-[0m[0;36m20[0m[0;34m,[0m [0;34m-[0m[0;36m20[0m[0;34m,[0m [0;34m-[0m[0;36m20[0m[0;34m,[0m [0;34m-[0m[0;36m18[0m[0;34m,[0m [0;34m-[0m[0;36m14[0m[0;34m,[0m [0;34m-[0m[0;36m22[0m[0;34m,[0m [0;34m-[0m[0;36m20[0m[0;34m,[0m [0;34m-[0m[0;36m14[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, -14, -20, -22, -14, -18, -20, -20, -20, -20, -18, -14, -22,
-20, -14, 0])In [ ]python · cell 6
python
