Prediction network Parameters ---------- action_space_size: int action space num_blocks: int number of res blocks num_channels: int channels of hidden states reduced_channels_value: int channels of value
(
self,
action_space_size,
num_blocks,
num_channels,
reduced_channels_value,
reduced_channels_policy,
fc_value_layers,
fc_policy_layers,
full_support_size,
block_output_size_value,
block_output_size_policy,
momentum=0.1,
init_zero=False,
)
| 291 | # predict the value and policy given hidden states |
| 292 | class PredictionNetwork(nn.Module): |
| 293 | def __init__( |
| 294 | self, |
| 295 | action_space_size, |
| 296 | num_blocks, |
| 297 | num_channels, |
| 298 | reduced_channels_value, |
| 299 | reduced_channels_policy, |
| 300 | fc_value_layers, |
| 301 | fc_policy_layers, |
| 302 | full_support_size, |
| 303 | block_output_size_value, |
| 304 | block_output_size_policy, |
| 305 | momentum=0.1, |
| 306 | init_zero=False, |
| 307 | ): |
| 308 | """Prediction network |
| 309 | Parameters |
| 310 | ---------- |
| 311 | action_space_size: int |
| 312 | action space |
| 313 | num_blocks: int |
| 314 | number of res blocks |
| 315 | num_channels: int |
| 316 | channels of hidden states |
| 317 | reduced_channels_value: int |
| 318 | channels of value head |
| 319 | reduced_channels_policy: int |
| 320 | channels of policy head |
| 321 | fc_value_layers: list |
| 322 | hidden layers of the value prediction head (MLP head) |
| 323 | fc_policy_layers: list |
| 324 | hidden layers of the policy prediction head (MLP head) |
| 325 | full_support_size: int |
| 326 | dim of value output |
| 327 | block_output_size_value: int |
| 328 | dim of flatten hidden states |
| 329 | block_output_size_policy: int |
| 330 | dim of flatten hidden states |
| 331 | init_zero: bool |
| 332 | True -> zero initialization for the last layer of value/policy mlp |
| 333 | """ |
| 334 | super().__init__() |
| 335 | self.resblocks = nn.ModuleList( |
| 336 | [ResidualBlock(num_channels, num_channels, momentum=momentum) for _ in range(num_blocks)] |
| 337 | ) |
| 338 | |
| 339 | self.conv1x1_value = nn.Conv2d(num_channels, reduced_channels_value, 1) |
| 340 | self.conv1x1_policy = nn.Conv2d(num_channels, reduced_channels_policy, 1) |
| 341 | self.bn_value = nn.BatchNorm2d(reduced_channels_value, momentum=momentum) |
| 342 | self.bn_policy = nn.BatchNorm2d(reduced_channels_policy, momentum=momentum) |
| 343 | self.block_output_size_value = block_output_size_value |
| 344 | self.block_output_size_policy = block_output_size_policy |
| 345 | self.fc_value = mlp(self.block_output_size_value, fc_value_layers, full_support_size, init_zero=init_zero, momentum=momentum) |
| 346 | self.fc_policy = mlp(self.block_output_size_policy, fc_policy_layers, action_space_size, init_zero=init_zero, momentum=momentum) |
| 347 | |
| 348 | def forward(self, x): |
| 349 | for block in self.resblocks: |
nothing calls this directly
no test coverage detected