MCPcopy Create free account
hub / github.com/YeWR/EfficientZero / __init__

Method __init__

config/atari/model.py:293–346  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

291# predict the value and policy given hidden states
292class 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:

Callers

nothing calls this directly

Calls 3

ResidualBlockClass · 0.85
mlpFunction · 0.85
__init__Method · 0.45

Tested by

no test coverage detected