()
| 121 | |
| 122 | |
| 123 | def main(): |
| 124 | env = gym.make('CartPole-v0') |
| 125 | ft = FeatureTransformer(env) |
| 126 | model = Model(env, ft) |
| 127 | gamma = 0.99 |
| 128 | |
| 129 | if 'monitor' in sys.argv: |
| 130 | filename = os.path.basename(__file__).split('.')[0] |
| 131 | monitor_dir = './' + filename + '_' + str(datetime.now()) |
| 132 | env = wrappers.Monitor(env, monitor_dir) |
| 133 | |
| 134 | |
| 135 | N = 500 |
| 136 | totalrewards = np.empty(N) |
| 137 | costs = np.empty(N) |
| 138 | for n in range(N): |
| 139 | eps = 1.0/np.sqrt(n+1) |
| 140 | totalreward = play_one(env, model, eps, gamma) |
| 141 | totalrewards[n] = totalreward |
| 142 | if n % 100 == 0: |
| 143 | print("episode:", n, "total reward:", totalreward, "eps:", eps, "avg reward (last 100):", totalrewards[max(0, n-100):(n+1)].mean()) |
| 144 | |
| 145 | print("avg reward for last 100 episodes:", totalrewards[-100:].mean()) |
| 146 | print("total steps:", totalrewards.sum()) |
| 147 | |
| 148 | plt.plot(totalrewards) |
| 149 | plt.title("Rewards") |
| 150 | plt.show() |
| 151 | |
| 152 | plot_running_avg(totalrewards) |
| 153 | |
| 154 | |
| 155 | if __name__ == '__main__': |
no test coverage detected