(self, state_dim, action_dim, max_action)
| 192 | |
| 193 | class SAC(): |
| 194 | def __init__(self, state_dim, action_dim, max_action): |
| 195 | super(SAC, self).__init__() |
| 196 | |
| 197 | parser = get_parameters() |
| 198 | self.args = parser.parse_args() |
| 199 | # Set seeds |
| 200 | torch.manual_seed(self.args.seed) |
| 201 | np.random.seed(self.args.seed) |
| 202 | |
| 203 | self.state_dim = state_dim |
| 204 | self.action_dim = action_dim |
| 205 | self.max_action = max_action |
| 206 | |
| 207 | self.policy_net = Actor(self.state_dim, self.action_dim).to(device) |
| 208 | self.value_net = Critic(self.state_dim).to(device) |
| 209 | self.Target_value_net = Critic(self.state_dim).to(device) |
| 210 | self.Q_net1 = Q(self.state_dim, self.action_dim).to(device) |
| 211 | self.Q_net2 = Q(self.state_dim, self.action_dim).to(device) |
| 212 | |
| 213 | self.policy_optimizer = optim.Adam(self.policy_net.parameters(), lr=self.args.learning_rate) |
| 214 | self.value_optimizer = optim.Adam(self.value_net.parameters(), lr=self.args.learning_rate) |
| 215 | self.Q1_optimizer = optim.Adam(self.Q_net1.parameters(), lr=self.args.learning_rate) |
| 216 | self.Q2_optimizer = optim.Adam(self.Q_net2.parameters(), lr=self.args.learning_rate) |
| 217 | |
| 218 | self.num_training = 1 |
| 219 | self.initial_tem = self.args.initial_tem |
| 220 | self.tem_decay_rate = self.args.tem_decay_rate |
| 221 | |
| 222 | self.last_map_level = 0 |
| 223 | self.num_training_map_level = 0 |
| 224 | |
| 225 | log_dir = './runs/' + current_time |
| 226 | self.writer = SummaryWriter(log_dir=log_dir) |
| 227 | |
| 228 | # calculate the loss |
| 229 | self.value_criterion = nn.MSELoss() |
| 230 | self.Q1_criterion = nn.MSELoss() |
| 231 | self.Q2_criterion = nn.MSELoss() |
| 232 | |
| 233 | # copy the weight of value_net to Target_value_net |
| 234 | for target_param, param in zip(self.Target_value_net.parameters(), self.value_net.parameters()): |
| 235 | target_param.data.copy_(param.data) |
| 236 | |
| 237 | def select_action(self, state): |
| 238 | state = torch.FloatTensor(state).to(device) |
nothing calls this directly
no test coverage detected