q value iteration

This commit is contained in:
Arjun Patel
2019-03-07 13:21:28 -08:00
parent 172b7c7dd3
commit b177b47969
2 changed files with 57 additions and 11 deletions
+1
View File
@@ -0,0 +1 @@
__pycache__
+53 -8
View File
@@ -26,11 +26,13 @@
# Pieter Abbeel (pabbeel@cs.berkeley.edu).
import mdp, util
import mdp
import util
from learningAgents import ValueEstimationAgent
import collections
class ValueIterationAgent(ValueEstimationAgent):
"""
* Please read learningAgents.py before reading this.*
@@ -40,7 +42,8 @@ class ValueIterationAgent(ValueEstimationAgent):
for a given number of iterations using the supplied
discount factor.
"""
def __init__(self, mdp, discount = 0.9, iterations = 100):
def __init__(self, mdp, discount=0.9, iterations=100):
"""
Your value iteration agent should take an mdp on
construction, run the indicated number of iterations
@@ -62,7 +65,26 @@ class ValueIterationAgent(ValueEstimationAgent):
def runValueIteration(self):
# Write value iteration code here
"*** YOUR CODE HERE ***"
INF, NEG_INF = float("inf"), -float("inf")
# Run through iterations
for i in range(self.iterations):
# copy function defined?
policy = self.values.copy()
# MDP states
mdp_states = self.mdp.getStates()
for curr_state in mdp_states:
# curr state is exit
if not self.mdp.isTerminal(curr_state):
options_actions = self.mdp.getPossibleActions(curr_state)
optimal = max([self.getQValue(curr_state, x)
for x in options_actions])
# add optimal to the policy
policy[curr_state] = optimal
# Update the new best policy
self.values = policy
def getValue(self, state):
"""
@@ -70,14 +92,22 @@ class ValueIterationAgent(ValueEstimationAgent):
"""
return self.values[state]
def computeQValueFromValues(self, state, action):
"""
Compute the Q-value of action in state from the
value function stored in self.values.
"""
"*** YOUR CODE HERE ***"
util.raiseNotDefined()
curr_val = 0
possible = self.mdp.getTransitionStatesAndProbs(state, action)
for new_state, prob in possible:
r = self.mdp.getReward(state, action, new_state)
val = self.values[new_state]
curr_val = curr_val + prob * ((self.discount * val) + r)
return curr_val
def computeActionFromValues(self, state):
"""
@@ -89,7 +119,19 @@ class ValueIterationAgent(ValueEstimationAgent):
terminal state, you should return None.
"""
"*** YOUR CODE HERE ***"
util.raiseNotDefined()
# end iteration
if self.mdp.isTerminal(state):
return None
curr_val, optimal_action = -float("inf"), ''
for action in self.mdp.getPossibleActions(state):
curr_qval = self.computeQValueFromValues(state, action)
# update if better
if curr_qval >= curr_val:
curr_val = curr_qval
optimal_action = action
return optimal_action
def getPolicy(self, state):
return self.computeActionFromValues(state)
@@ -101,6 +143,7 @@ class ValueIterationAgent(ValueEstimationAgent):
def getQValue(self, state, action):
return self.computeQValueFromValues(state, action)
class AsynchronousValueIterationAgent(ValueIterationAgent):
"""
* Please read learningAgents.py before reading this.*
@@ -110,7 +153,8 @@ class AsynchronousValueIterationAgent(ValueIterationAgent):
for a given number of iterations using the supplied
discount factor.
"""
def __init__(self, mdp, discount = 0.9, iterations = 1000):
def __init__(self, mdp, discount=0.9, iterations=1000):
"""
Your cyclic value iteration agent should take an mdp on
construction, run the indicated number of iterations,
@@ -131,6 +175,7 @@ class AsynchronousValueIterationAgent(ValueIterationAgent):
def runValueIteration(self):
"*** YOUR CODE HERE ***"
class PrioritizedSweepingValueIterationAgent(AsynchronousValueIterationAgent):
"""
* Please read learningAgents.py before reading this.*
@@ -139,7 +184,8 @@ class PrioritizedSweepingValueIterationAgent(AsynchronousValueIterationAgent):
(see mdp.py) on initialization and runs prioritized sweeping value iteration
for a given number of iterations using the supplied parameters.
"""
def __init__(self, mdp, discount = 0.9, iterations = 100, theta = 1e-5):
def __init__(self, mdp, discount=0.9, iterations=100, theta=1e-5):
"""
Your prioritized sweeping value iteration agent should take an mdp on
construction, run the indicated number of iterations,
@@ -150,4 +196,3 @@ class PrioritizedSweepingValueIterationAgent(AsynchronousValueIterationAgent):
def runValueIteration(self):
"*** YOUR CODE HERE ***"