| 173 | |
| 174 | |
| 175 | class CVTNet(nn.Module): |
| 176 | def __init__(self, channels=5, use_transformer = True): |
| 177 | super(CVTNet, self).__init__() |
| 178 | |
| 179 | self.use_transformer = use_transformer |
| 180 | |
| 181 | self.featureExtracter_RI = featureExtracter_RI_BEV(channels=channels, use_transformer=use_transformer) |
| 182 | self.featureExtracter_BEV = featureExtracter_RI_BEV(channels=channels, use_transformer=use_transformer) |
| 183 | |
| 184 | self.relu = nn.ReLU(inplace=True) |
| 185 | |
| 186 | d_model = 256 |
| 187 | heads = 4 |
| 188 | dropout = 0. |
| 189 | |
| 190 | self.convLast2 = nn.Conv2d(256, 256, kernel_size=(1,1), stride=(1,1), bias=False) |
| 191 | self.sigmoid = nn.Sigmoid() |
| 192 | self.softmax = nn.Softmax() |
| 193 | |
| 194 | self.net_vlad = NetVLADLoupe(feature_size=512, max_samples=1800, cluster_size=64, |
| 195 | output_dim=256, gating=True, add_batch_norm=False, |
| 196 | is_training=True) |
| 197 | self.net_vlad_ri = NetVLADLoupe(feature_size=256, max_samples=900, cluster_size=64, |
| 198 | output_dim=256, gating=True, add_batch_norm=False, |
| 199 | is_training=True) |
| 200 | self.net_vlad_bev = NetVLADLoupe(feature_size=256, max_samples=900, cluster_size=64, |
| 201 | output_dim=256, gating=True, add_batch_norm=False, |
| 202 | is_training=True) |
| 203 | self.norm_1 = Norm(d_model) |
| 204 | self.norm_2 = Norm(d_model) |
| 205 | self.norm_3 = Norm(d_model) |
| 206 | self.norm_2_ext = Norm(d_model) |
| 207 | self.norm_3_ext = Norm(d_model) |
| 208 | |
| 209 | self.attn1 = MultiHeadAttention(heads, d_model, dropout=dropout) |
| 210 | self.attn2 = MultiHeadAttention(heads, d_model, dropout=dropout) |
| 211 | |
| 212 | self.ff1 = FeedForward(d_model, dropout=dropout) |
| 213 | self.ff2 = FeedForward(d_model, dropout=dropout) |
| 214 | |
| 215 | self.attn1_ext = MultiHeadAttention(heads, d_model, dropout=dropout) |
| 216 | self.attn2_ext = MultiHeadAttention(heads, d_model, dropout=dropout) |
| 217 | |
| 218 | self.ff1_ext = FeedForward(d_model, dropout=dropout) |
| 219 | self.ff2_ext = FeedForward(d_model, dropout=dropout) |
| 220 | |
| 221 | |
| 222 | def forward(self, x_ri_bev): |
| 223 | x_ri = x_ri_bev[:, 0:5, :, :] |
| 224 | x_bev = x_ri_bev[:, 5:10, :, :] |
| 225 | |
| 226 | feature_ri = self.featureExtracter_RI(x_ri) |
| 227 | feature_bev = self.featureExtracter_BEV(x_bev) |
| 228 | |
| 229 | feature_ri = feature_ri.squeeze(-1) |
| 230 | feature_bev = feature_bev.squeeze(-1) |
| 231 | feature_ri = feature_ri.permute(0, 2, 1) |
| 232 | feature_bev = feature_bev.permute(0, 2, 1) |