# BATCH_SIZE is the number of transitions sampled from the replay buffer # GAMMA is the discount factor as mentioned in the previous section # EPS_START is the starting value of epsilon # EPS_END is the final value of epsilon # EPS_DECAY controls the rate of exponential decay of epsilon, higher means a slower decay # TAU is the update rate of the target network # LR is the learning rate of the ``AdamW`` optimizer BATCH_SIZE = 128 GAMMA = 0.99 EPS_START = 0.9 EPS_END = 0.01 EPS_DECAY = 2500 TAU = 0.005 LR = 3e-4 # Get number of actions from gym action space n_actions = env.action_space.n # Get the number of state observations [state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), info = env.reset() n_observations = len([state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) policy_net = [DQN](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module "torch.nn.Module")(n_observations, n_actions).to([device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device")) target_net = [DQN](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module "torch.nn.Module")(n_observations, n_actions).to([device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device")) [target_net.load_state_dict](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.load_state_dict "torch.nn.Module.load_state_dict")([policy_net.state_dict](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.state_dict "torch.nn.Module.state_dict")()) [optimizer](https://docs.pytorch.org/docs/stable/generated/torch.optim.AdamW.html#torch.optim.AdamW "torch.optim.AdamW") = [optim.AdamW](https://docs.pytorch.org/docs/stable/generated/torch.optim.AdamW.html#torch.optim.AdamW "torch.optim.AdamW")([policy_net.parameters](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.parameters "torch.nn.Module.parameters")(), lr=LR, amsgrad=True) memory = ReplayMemory(10000) steps_done = 0 def select_action([state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")): global steps_done sample = random.random() eps_threshold = EPS_END + (EPS_START - EPS_END) * \ math.exp(-1. * steps_done / EPS_DECAY) steps_done += 1 if sample > eps_threshold: with [torch.no_grad](https://docs.pytorch.org/docs/stable/generated/torch.no_grad.html#torch.no_grad "torch.no_grad")(): # t.max(1) will return the largest column value of each row. # second column on max result is index of where max element was # found, so we pick action with the larger expected reward. return policy_net([state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")).max(1).indices.view(1, 1) else: return [torch.tensor](https://docs.pytorch.org/docs/stable/generated/torch.tensor.html#torch.tensor "torch.tensor")([[env.action_space.sample()]], [device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device")=[device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device"), dtype=[torch.long](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.dtype "torch.dtype")) episode_durations = [] def plot_durations(show_result=False): plt.figure(1) durations_t = [torch.tensor](https://docs.pytorch.org/docs/stable/generated/torch.tensor.html#torch.tensor "torch.tensor")(episode_durations, dtype=[torch.float](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.dtype "torch.dtype")) if show_result: plt.title('Result') else: plt.clf() plt.title('Training...') plt.xlabel('Episode') plt.ylabel('Duration') plt.plot(durations_t.numpy()) # Take 100 episode averages and plot them too if len(durations_t) >= 100: means = durations_t.unfold(0, 100, 1).mean(1).view(-1) means = [torch.cat](https://docs.pytorch.org/docs/stable/generated/torch.cat.html#torch.cat "torch.cat")(([torch.zeros](https://docs.pytorch.org/docs/stable/generated/torch.zeros.html#torch.zeros "torch.zeros")(99), means)) plt.plot(means.numpy()) plt.pause(0.001) # pause a bit so that plots are updated if is_ipython: if not show_result: display.display(plt.gcf()) display.clear_output(wait=True) else: display.display(plt.gcf())