if [torch.cuda.is_available](https://docs.pytorch.org/docs/stable/generated/torch.cuda.is_available.html#torch.cuda.is_available "torch.cuda.is_available")() or [torch.backends.mps.is_available](https://docs.pytorch.org/docs/stable/backends.html#torch.backends.mps.is_available "torch.backends.mps.is_available")(): num_episodes = 600 else: num_episodes = 50 for i_episode in range(num_episodes): # Initialize the environment and get its state [state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), info = env.reset() [state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [torch.tensor](https://docs.pytorch.org/docs/stable/generated/torch.tensor.html#torch.tensor "torch.tensor")([state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), dtype=[torch.float32](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.dtype "torch.dtype"), [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")).unsqueeze(0) for t in count(): [action](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = select_action([state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) observation, [reward](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), terminated, truncated, _ = env.step([action](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").item()) [reward](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [torch.tensor](https://docs.pytorch.org/docs/stable/generated/torch.tensor.html#torch.tensor "torch.tensor")([[reward](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")], [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")) done = terminated or truncated if terminated: [next_state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = None else: [next_state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [torch.tensor](https://docs.pytorch.org/docs/stable/generated/torch.tensor.html#torch.tensor "torch.tensor")(observation, dtype=[torch.float32](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.dtype "torch.dtype"), [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")).unsqueeze(0) # Store the transition in memory memory.push([state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), [action](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), [next_state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), [reward](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) # Move to the next state [state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [next_state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") # Perform one step of the optimization (on the policy network) optimize_model() # Soft update of the target network's weights # θ′ ← τ θ + (1 −τ )θ′ target_net_state_dict = [target_net.state_dict](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.state_dict "torch.nn.Module.state_dict")() policy_net_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")() for key in policy_net_state_dict: target_net_state_dict[key] = policy_net_state_dict[key]*TAU + target_net_state_dict[key]*(1-TAU) [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")(target_net_state_dict) if done: episode_durations.append(t + 1) plot_durations() break print('Complete') plot_durations(show_result=True) plt.ioff() plt.show()