- improved

git-svn-id: http://moon:8086/svn/projects/RL-lab@343 fda53097-d464-4ada-af97-ba876c37ca34
This commit is contained in:
2020-01-02 21:26:07 +00:00
parent 7d06ea0d23
commit 7d00a4a376
+4 -5
View File
@@ -71,9 +71,9 @@ def run(pos, N_trials, learning_rate, penalty_factor=0.99, forgetting_factor=0.9
# forget(forgetting_factor) # forget(forgetting_factor)
P = P_door[pos] Pn = P_door[pos]/np.sum(P_door[pos])
Pn = P/np.sum(P)
P_door[pos,:] *= forgetting_factor
while True: while True:
fail = False fail = False
door = choose_door(Pn) door = choose_door(Pn)
@@ -100,8 +100,7 @@ def run(pos, N_trials, learning_rate, penalty_factor=0.99, forgetting_factor=0.9
P_door[pos][door] *= penalty_factor P_door[pos][door] *= penalty_factor
continue continue
else: else:
P[door] = P[door] + learning_rate P_door[pos][door] = P_door[pos][door] + learning_rate
P_door[pos] = P
break break
@@ -112,7 +111,7 @@ def run(pos, N_trials, learning_rate, penalty_factor=0.99, forgetting_factor=0.9
return moves_needed, lab_visited return moves_needed, lab_visited
for i in range(0, 1000): for i in range(0, 100):
moves_needed, lab_visited = run(pos=(0,0), N_trials=1000, learning_rate=0.5) moves_needed, lab_visited = run(pos=(0,0), N_trials=1000, learning_rate=0.5)
print ('Visited map after {} moves:'.format(moves_needed)) print ('Visited map after {} moves:'.format(moves_needed))
# print (lab_visited) # print (lab_visited)