1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93
| class DeepFusion(nn.Module): """ 深度融合(特征级) 多阶段跨模态特征交互 """ def __init__(self, dim=256): super().__init__() self.rgb_enc = FeatureEncoder(3, dim) self.lidar_enc = FeatureEncoder(1, dim) self.cross_attn = CrossModalAttention(dim) self.det_head = DetectionHead(dim) def forward(self, rgb, lidar): """ Args: rgb: RGB图像 lidar: LiDAR数据 Returns: detections: 检测结果 """ rgb_feat = self.rgb_enc(rgb) lidar_feat = self.lidar_enc(lidar) rgb_enhanced, lidar_enhanced = self.cross_attn(rgb_feat, lidar_feat) fused = rgb_enhanced + lidar_enhanced detections = self.det_head(fused) return detections
class FeatureEncoder(nn.Module): """特征编码器""" def __init__(self, in_channels, dim): super().__init__() self.encoder = nn.Sequential( nn.Conv2d(in_channels, 64, 7, stride=2, padding=3), nn.ReLU(), nn.Conv2d(64, dim, 3, stride=2, padding=1), nn.ReLU() ) def forward(self, x): return self.encoder(x)
class CrossModalAttention(nn.Module): """跨模态注意力""" def __init__(self, dim): super().__init__() self.attn = nn.MultiheadAttention(dim, num_heads=8) def forward(self, feat1, feat2): B, C, H, W = feat1.shape feat1_flat = feat1.flatten(2).transpose(0, 1) feat2_flat = feat2.flatten(2).transpose(0, 1) enh1, _ = self.attn(feat1_flat, feat2_flat, feat2_flat) enh2, _ = self.attn(feat2_flat, feat1_flat, feat1_flat) enh1 = enh1.transpose(0, 1).view(B, C, H, W) enh2 = enh2.transpose(0, 1).view(B, C, H, W) return enh1, enh2
class DetectionHead(nn.Module): """检测头""" def __init__(self, dim): super().__init__() self.head = nn.Conv2d(dim, 6, 1) def forward(self, x): return self.head(x)
|