| 13 | ''' |
| 14 | |
| 15 | def __init__(self, observation_space, features_dim=256, states_neurons=[256]): |
| 16 | super().__init__() |
| 17 | self.features_dim = features_dim |
| 18 | |
| 19 | n_input_channels = observation_space['birdview'].shape[0] |
| 20 | |
| 21 | self.cnn = nn.Sequential( |
| 22 | nn.Conv2d(n_input_channels, 8, kernel_size=5, stride=2), |
| 23 | nn.ReLU(), |
| 24 | nn.Conv2d(8, 16, kernel_size=5, stride=2), |
| 25 | nn.ReLU(), |
| 26 | nn.Conv2d(16, 32, kernel_size=5, stride=2), |
| 27 | nn.ReLU(), |
| 28 | nn.Conv2d(32, 64, kernel_size=3, stride=2), |
| 29 | nn.ReLU(), |
| 30 | nn.Conv2d(64, 128, kernel_size=3, stride=2), |
| 31 | nn.ReLU(), |
| 32 | nn.Conv2d(128, 256, kernel_size=3, stride=1), |
| 33 | nn.ReLU(), |
| 34 | nn.Flatten(), |
| 35 | ) |
| 36 | # Compute shape by doing one forward pass |
| 37 | with th.no_grad(): |
| 38 | n_flatten = self.cnn(th.as_tensor(observation_space['birdview'].sample()[None]).float()).shape[1] |
| 39 | |
| 40 | self.linear = nn.Sequential(nn.Linear(n_flatten+states_neurons[-1], 512), nn.ReLU(), |
| 41 | nn.Linear(512, features_dim), nn.ReLU()) |
| 42 | |
| 43 | states_neurons = [observation_space['state'].shape[0]] + states_neurons |
| 44 | self.state_linear = [] |
| 45 | for i in range(len(states_neurons)-1): |
| 46 | self.state_linear.append(nn.Linear(states_neurons[i], states_neurons[i+1])) |
| 47 | self.state_linear.append(nn.ReLU()) |
| 48 | self.state_linear = nn.Sequential(*self.state_linear) |
| 49 | |
| 50 | self.apply(self._weights_init) |
| 51 | |
| 52 | @staticmethod |
| 53 | def _weights_init(m): |