(cls, path)
| 220 | |
| 221 | @classmethod |
| 222 | def load(cls, path): |
| 223 | if th.cuda.is_available(): |
| 224 | device = 'cuda' |
| 225 | else: |
| 226 | device = 'cpu' |
| 227 | saved_variables = th.load(path, map_location=device) |
| 228 | # Create policy object |
| 229 | model = cls(**saved_variables['policy_init_kwargs']) |
| 230 | |
| 231 | state_dict = saved_variables['policy_state_dict'] |
| 232 | for index in [2, 4, 6, 8, 10]: |
| 233 | state_dict["features_extractor.cnn."+str(index//2)+".weight"] = state_dict.pop("features_extractor.cnn."+str(index)+".weight") |
| 234 | state_dict["features_extractor.cnn."+str(index//2)+".bias"] = state_dict.pop("features_extractor.cnn."+str(index)+".bias") |
| 235 | #OLD!!! odict_keys(['features_extractor.cnn.0.weight', 'features_extractor.cnn.0.bias', 'features_extractor.cnn.2.weight', 'features_extractor.cnn.2.bias', 'features_extractor.cnn.4.weight', 'features_extractor.cnn.4.bias', 'features_extractor.cnn.6.weight', 'features_extractor.cnn.6.bias', 'features_extractor.cnn.8.weight', 'features_extractor.cnn.8.bias', 'features_extractor.cnn.10.weight', 'features_extractor.cnn.10.bias', 'features_extractor.linear.0.weight', 'features_extractor.linear.0.bias', 'features_extractor.linear.2.weight', 'features_extractor.linear.2.bias', 'features_extractor.state_linear.0.weight', 'features_extractor.state_linear.0.bias', 'features_extractor.state_linear.2.weight', 'features_extractor.state_linear.2.bias', 'policy_head.0.weight', 'policy_head.0.bias', 'policy_head.2.weight', 'policy_head.2.bias', 'dist_mu.0.weight', 'dist_mu.0.bias', 'dist_sigma.0.weight', 'dist_sigma.0.bias', 'value_head.0.weight', 'value_head.0.bias', 'value_head.2.weight', 'value_head.2.bias', 'value_head.4.weight', 'value_head.4.bias']) |
| 236 | # odict_keys(['features_extractor.cnn.0.weight', 'features_extractor.cnn.0.bias', 'features_extractor.cnn.1.weight', 'features_extractor.cnn.1.bias', 'features_extractor.cnn.2.weight', 'features_extractor.cnn.2.bias', 'features_extractor.cnn.3.weight', 'features_extractor.cnn.3.bias', 'features_extractor.cnn.4.weight', 'features_extractor.cnn.4.bias', 'features_extractor.cnn.5.weight', 'features_extractor.cnn.5.bias', 'features_extractor.linear.0.weight', 'features_extractor.linear.0.bias', 'features_extractor.linear.2.weight', 'features_extractor.linear.2.bias', 'features_extractor.state_linear.0.weight', 'features_extractor.state_linear.0.bias', 'features_extractor.state_linear.2.weight', 'features_extractor.state_linear.2.bias', 'policy_head.0.weight', 'policy_head.0.bias', 'policy_head.2.weight', 'policy_head.2.bias', 'dist_mu.0.weight', 'dist_mu.0.bias', 'dist_sigma.0.weight', 'dist_sigma.0.bias', 'value_head.0.weight', 'value_head.0.bias', 'value_head.2.weight', 'value_head.2.bias', 'value_head.4.weight', 'value_head.4.bias']) |
| 237 | # Load weights |
| 238 | model.load_state_dict(saved_variables['policy_state_dict']) |
| 239 | model.to(device) |
| 240 | return model, saved_variables['train_init_kwargs'] |
| 241 | |
| 242 | @staticmethod |
| 243 | def init_weights(module: nn.Module, gain: float = 1) -> None: |
no outgoing calls
no test coverage detected