Skip to content

Commit aeed84f

Browse files
committed
Handle set actlist in MDP.__init__ (#1259)
MDP.__init__ only assigned self.actlist for list and dict inputs; a set fell through both branches, leaving the attribute unset so MDP.actions() raised AttributeError (and value/policy iteration crashed). Add a set branch that stores the actions as a list, keeping them indexable for random.choice in policy_iteration. Mirror the fix in mdp4e.py and extend the existing set-actlist test to exercise actions().
1 parent ca11f95 commit aeed84f

3 files changed

Lines changed: 12 additions & 0 deletions

File tree

‎mdp.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,10 @@ def __init__(self, init, actlist, terminals, transitions=None, reward=None, stat
4242
# if actlist is a dict, different actions for each state
4343
self.actlist = actlist
4444

45+
elif isinstance(actlist, set):
46+
# if actlist is a set, convert it to a list so actions are indexable
47+
self.actlist = list(actlist)
48+
4549
self.terminals = terminals
4650
self.transitions = transitions or {}
4751
if not self.transitions:

‎mdp4e.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,10 @@ def __init__(self, init, actlist, terminals, transitions=None, reward=None, stat
4242
# if actlist is a dict, different actions for each state
4343
self.actlist = actlist
4444

45+
elif isinstance(actlist, set):
46+
# if actlist is a set, convert it to a list so actions are indexable
47+
self.actlist = list(actlist)
48+
4549
self.terminals = terminals
4650
self.transitions = transitions or {}
4751
if not self.transitions:

‎tests/test_mdp.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -126,6 +126,10 @@ def test_transition_model():
126126
assert mdp.T("b", "plan2") == [(0.6, 'a'), (0.2, 'b'), (0.1, 'c'), (0.1, 'd')]
127127
assert mdp.T("c", "plan1") == [(0.3, 'a'), (0.5, 'b'), (0.1, 'c'), (0.1, 'd')]
128128

129+
# a set actlist must be accepted and yield indexable per-state actions
130+
assert set(mdp.actions("a")) == {"plan1", "plan2", "plan3"}
131+
assert mdp.actions("d") == [None]
132+
129133

130134
def test_pomdp_value_iteration():
131135
t_prob = [[[0.65, 0.35], [0.65, 0.35]], [[0.65, 0.35], [0.65, 0.35]], [[1.0, 0.0], [0.0, 1.0]]]

0 commit comments

Comments
 (0)