def optimize_model(): if len(memory) < BATCH_SIZE: return transitions = memory.sample(BATCH_SIZE) # Transpose the batch (see https://stackoverflow.com/a/19343/3343043 for # detailed explanation). This converts batch-array of Transitions # to Transition of batch-arrays. batch = Transition(*zip(*transitions)) # Compute a mask of non-final states and concatenate the batch elements # (a final state would've been the one after which simulation ended) non_final_mask = [torch.tensor](https://docs.pytorch.org/docs/stable/generated/torch.tensor.html#torch.tensor "torch.tensor")(tuple(map(lambda s: s is not None, batch.[next_state](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"), dtype=[torch.bool](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.dtype "torch.dtype")) non_final_next_states = [torch.cat](https://docs.pytorch.org/docs/stable/generated/torch.cat.html#torch.cat "torch.cat")([s for s in batch.[next_state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") if s is not None]) state_batch = [torch.cat](https://docs.pytorch.org/docs/stable/generated/torch.cat.html#torch.cat "torch.cat")(batch.[state](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) action_batch = [torch.cat](https://docs.pytorch.org/docs/stable/generated/torch.cat.html#torch.cat "torch.cat")(batch.[action](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) reward_batch = [torch.cat](https://docs.pytorch.org/docs/stable/generated/torch.cat.html#torch.cat "torch.cat")(batch.[reward](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) # Compute Q(s_t, a) - the model computes Q(s_t), then we select the # columns of actions taken. These are the actions which would've been taken # for each batch state according to policy_net state_action_values = policy_net(state_batch).gather(1, action_batch) # Compute V(s_{t+1}) for all next states. # Expected values of actions for non_final_next_states are computed based # on the "older" target_net; selecting their best reward with max(1).values # This is merged based on the mask, such that we'll have either the expected # state value or 0 in case the state was final. next_state_values = [torch.zeros](https://docs.pytorch.org/docs/stable/generated/torch.zeros.html#torch.zeros "torch.zeros")(BATCH_SIZE, [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")) with [torch.no_grad](https://docs.pytorch.org/docs/stable/generated/torch.no_grad.html#torch.no_grad "torch.no_grad")(): next_state_values[non_final_mask] = target_net(non_final_next_states).max(1).values # Compute the expected Q values expected_state_action_values = (next_state_values * GAMMA) + reward_batch # Compute Huber loss criterion = [nn.SmoothL1Loss](https://docs.pytorch.org/docs/stable/generated/torch.nn.SmoothL1Loss.html#torch.nn.SmoothL1Loss "torch.nn.SmoothL1Loss")() loss = criterion(state_action_values, expected_state_action_values.unsqueeze(1)) # Optimize the model [optimizer.zero_grad](https://docs.pytorch.org/docs/stable/generated/torch.optim.AdamW.html#torch.optim.AdamW.zero_grad "torch.optim.AdamW.zero_grad")() loss.backward() # In-place gradient clipping [torch.nn.utils.clip_grad_value_](https://docs.pytorch.org/docs/stable/generated/torch.nn.utils.clip_grad_value_.html#torch.nn.utils.clip_grad_value_ "torch.nn.utils.clip_grad_value_")([policy_net.parameters](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.parameters "torch.nn.Module.parameters")(), 100) [optimizer.step](https://docs.pytorch.org/docs/stable/generated/torch.optim.AdamW.html#torch.optim.AdamW.step "torch.optim.AdamW.step")()