MCPcopy Create free account
hub / github.com/BIT-MJY/CVTNet / CVTNet

Class CVTNet

modules/cvtnet.py:175–279  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

173
174
175class 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)

Callers 4

__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90

Calls

no outgoing calls

Tested by 2

__init__Method · 0.72
__init__Method · 0.72