MCPcopy Create free account
hub / github.com/OpenDriveLab/DriveAdapter / __init__

Method __init__

roach/models/torch_layers.py:15–66  ·  view source on GitHub ↗
(self, observation_space, features_dim=256, states_neurons=[256])

Source from the content-addressed store, hash-verified

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.ModuleList([
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 self.relu = nn.ReLU()
37 # self.cnn = nn.ModuleList(
38 # nn.Conv2d(n_input_channels, 8, kernel_size=5, stride=2),
39 # nn.ReLU(),
40 # nn.Conv2d(8, 16, kernel_size=5, stride=2),
41 # nn.ReLU(),
42 # nn.Conv2d(16, 32, kernel_size=5, stride=2),
43 # nn.ReLU(),
44 # nn.Conv2d(32, 64, kernel_size=3, stride=2),
45 # nn.ReLU(),
46 # nn.Conv2d(64, 128, kernel_size=3, stride=2),
47 # nn.ReLU(),
48 # nn.Conv2d(128, 256, kernel_size=3, stride=1),
49 # nn.ReLU(),
50 # #nn.Flatten(),
51 # )
52 # Compute shape by doing one forward pass
53 #with th.no_grad():
54 #n_flatten = self.cnn(th.as_tensor(observation_space['birdview'].sample()[None]).float()).flatten(start_dim=1).shape[1]
55 n_flatten = 1024
56 self.linear = nn.Sequential(nn.Linear(n_flatten+states_neurons[-1], 512), nn.ReLU(),
57 nn.Linear(512, features_dim), nn.ReLU())
58
59 states_neurons = [observation_space['state'].shape[0]] + states_neurons
60 self.state_linear = []
61 for i in range(len(states_neurons)-1):
62 self.state_linear.append(nn.Linear(states_neurons[i], states_neurons[i+1]))
63 self.state_linear.append(nn.ReLU())
64 self.state_linear = nn.Sequential(*self.state_linear)
65
66 self.apply(self._weights_init)
67
68 @staticmethod
69 def _weights_init(m):

Callers 1

__init__Method · 0.45

Calls

no outgoing calls

Tested by

no test coverage detected