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 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139
| import torch import torch.nn as nn from typing import Tuple
class AIDEBaseline(nn.Module): """ AIDE数据集基线模型架构 支持三种建模模式: 1. 2D Pattern: 单帧CNN 2. 2D+Timing: CNN + Transformer Encoder 3. 3D Pattern: 时空模型 """ def __init__(self, pattern: str = "2d_timing", num_views: int = 4): super().__init__() self.pattern = pattern self.num_views = num_views if pattern == "2d": self.backbone = ResNet50(pretrained=True) self.head = nn.Linear(2048, num_classes) elif pattern == "2d_timing": self.backbone = ResNet50(pretrained=True) self.temporal_encoder = nn.TransformerEncoder( nn.TransformerEncoderLayer( d_model=2048, nhead=8, dim_feedforward=8192, dropout=0.1, batch_first=True ), num_layers=4 ) self.head = nn.Linear(2048, num_classes) elif pattern == "3d": self.backbone = TimeSFormer( img_size=224, num_frames=16, patch_size=16, num_classes=num_classes ) def forward(self, x: Tuple[torch.Tensor, ...]) -> torch.Tensor: """ Args: x: 多视角输入元组,每个 (B, C, T, H, W) front, left, right, inside Returns: logits: (B, num_classes) """ if self.pattern == "2d_timing": view_features = [] for view_idx in range(self.num_views): view_input = x[view_idx] B, C, T, H, W = view_input.shape frames = view_input.permute(0, 2, 1, 3, 4) frame_features = [] for t in range(T): feat = self.backbone(frames[:, t]) frame_features.append(feat) frame_features = torch.stack(frame_features, dim=1) encoded = self.temporal_encoder(frame_features) view_features.append(encoded.mean(dim=1)) fused = torch.stack(view_features, dim=1) fused = fused.mean(dim=1) return self.head(fused) elif self.pattern == "3d": inside_view = x[3] return self.backbone(inside_view)
class AdaptiveFusionModule(nn.Module): """ 自适应融合模块 根据场景动态加权不同模态: - 直行时:内部视角权重高 - 转弯时:外部视角权重高 - 危险时:所有视角均重要 """ def __init__(self, feature_dim: int = 2048, num_views: int = 4): super().__init__() self.attention = nn.MultiheadAttention( embed_dim=feature_dim, num_heads=8, batch_first=True ) self.view_weights = nn.Parameter(torch.ones(num_views)) def forward(self, view_features: torch.Tensor) -> torch.Tensor: """ Args: view_features: (B, num_views, feature_dim) Returns: fused: (B, feature_dim) """ weights = torch.softmax(self.view_weights, dim=0) weighted = view_features * weights.unsqueeze(0).unsqueeze(-1) fused, _ = self.attention(weighted, weighted, weighted) return fused.mean(dim=1)
if __name__ == "__main__": batch_size = 2 num_frames = 16 front = torch.randn(batch_size, 3, num_frames, 224, 224) left = torch.randn(batch_size, 3, num_frames, 224, 224) right = torch.randn(batch_size, 3, num_frames, 224, 224) inside = torch.randn(batch_size, 3, num_frames, 224, 224) model = AIDEBaseline(pattern="2d_timing", num_views=4) output = model((front, left, right, inside)) print(f"Output shape: {output.shape}") print(f"Model params: {sum(p.numel() for p in model.parameters()) / 1e6:.1f}M")
|