Created
February 11, 2020 15:14
-
-
Save ikatsov/a267ab5423abd93eeb2f5f89f8c3be40 to your computer and use it in GitHub Desktop.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| policy_net = PolicyNetworkDQN(2*T, len(price_grid)).to(device) | |
| target_net = PolicyNetworkDQN(2*T, len(price_grid)).to(device) | |
| optimizer = optim.AdamW(policy_net.parameters(), lr = 0.005) | |
| policy = AnnealedEpsGreedyPolicy() | |
| memory = ReplayMemory(10000) | |
| TARGET_UPDATE = 20 | |
| num_episodes = 1000 | |
| return_trace = [] | |
| p_trace = [] # Price schedules used in each episode | |
| for i_episode in range(num_episodes): | |
| state = env_intial_state() | |
| reward_trace = [] | |
| p = [] | |
| for t in range(T): | |
| # Select and perform an action | |
| with torch.no_grad(): | |
| q_values = policy_net(to_tensor(state)) | |
| action = policy.select_action(q_values.detach().numpy()) | |
| next_state, reward = env_step(t, state, action) | |
| # Store the transition in memory | |
| memory.push(to_tensor(state), | |
| to_tensor_long(action), | |
| to_tensor(next_state) if t != T - 1 else None, | |
| to_tensor([reward])) | |
| # Move to the next state | |
| state = next_state | |
| # Perform one step of the optimization | |
| update_model(memory, policy_net, target_net) | |
| reward_trace.append(reward) | |
| p.append(price_grid[action]) | |
| return_trace.append(sum(reward_trace)) | |
| p_trace.append(p) | |
| # Update the target network, copying all weights and biases in DQN | |
| if i_episode % TARGET_UPDATE == 0: | |
| target_net.load_state_dict(policy_net.state_dict()) |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment