| 114 | |
| 115 | |
| 116 | class NNBase(nn.Module): |
| 117 | def __init__(self, num_inputs, device, trainable=True, hidden_size=params.temporal_hidden_size): |
| 118 | super(NNBase, self).__init__() |
| 119 | self._feature_size = hidden_size |
| 120 | |
| 121 | if params.use_rnn is False or params.use_spatial_att is False: |
| 122 | init_ = lambda m: init(m, nn.init.orthogonal_, lambda x: nn.init.constant_(x, 0), |
| 123 | nn.init.calculate_gain('relu')) |
| 124 | |
| 125 | self.feature = nn.Sequential( |
| 126 | # input: 3*80*80 |
| 127 | init_(nn.Conv2d(num_inputs, 32, 8, stride=4)), |
| 128 | nn.LayerNorm([32, 19, 19]), |
| 129 | nn.ReLU(inplace=True), |
| 130 | # input: 32*19*19 |
| 131 | init_(nn.Conv2d(32, 64, 4, stride=2)), |
| 132 | nn.LayerNorm([64, 8, 8]), |
| 133 | nn.ReLU(inplace=True), |
| 134 | # input: 64*8*8 |
| 135 | init_(nn.Conv2d(64, 32, 3, stride=1)), |
| 136 | nn.LayerNorm([32, 6, 6]), |
| 137 | nn.ReLU(inplace=True), |
| 138 | # output: 32*6*6 |
| 139 | ).to(device) |
| 140 | |
| 141 | init_ = lambda m: init(m, |
| 142 | nn.init.orthogonal_, |
| 143 | lambda x: nn.init.constant_(x, 0)) |
| 144 | |
| 145 | self.conv_to_flat = nn.Sequential( |
| 146 | Flatten(), |
| 147 | init_(nn.Linear(32 * 6 * 6, self._feature_size)), |
| 148 | nn.LayerNorm([self._feature_size]), |
| 149 | nn.ReLU(inplace=True), |
| 150 | ).to(device) |
| 151 | |
| 152 | if params.use_rnn: |
| 153 | if params.use_relational_att: |
| 154 | # self.gru = RelationalGRU(input_size=hidden_size, hidden_dim=hidden_size, use_att=params.use_att).to( |
| 155 | # device) |
| 156 | self.gru = RelationalGRU(input_size=hidden_size, hidden_dim=hidden_size, use_att=False).to( |
| 157 | device) |
| 158 | elif params.use_spatial_att: |
| 159 | self.gru = SpatialAttGRU(input_dim=num_inputs).to(device) |
| 160 | else: |
| 161 | self.gru = nn.GRU(hidden_size, hidden_size).to(device) |
| 162 | |
| 163 | for name, param in self.gru.named_parameters(): |
| 164 | if 'bias' in name: |
| 165 | nn.init.constant_(param, 0) |
| 166 | elif 'weight' in name: |
| 167 | nn.init.orthogonal_(param) |
| 168 | |
| 169 | if trainable: |
| 170 | self.train() |
| 171 | else: |
| 172 | self.eval() |
| 173 | for p in self.parameters(): |