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 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178
| import torch import torch.nn as nn import torch.nn.functional as F
class FineGrainedInterFrameAttention(nn.Module): """ FIFA: 细粒度帧间注意力模块 论文: FIFA: Fine-grained Inter-frame Attention for Driver's Video Gaze Estimation (CVPR 2025) 核心思想: 在相邻视频帧之间建立细粒度注意力, 捕捉驾驶员视线的微妙时序变化 """ def __init__(self, feat_dim=512, num_heads=8, window_size=5): super().__init__() self.feat_dim = feat_dim self.num_heads = num_heads self.window_size = window_size self.frame_encoder = nn.Sequential( nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3), nn.BatchNorm2d(64), nn.ReLU(inplace=True), nn.MaxPool2d(kernel_size=3, stride=2, padding=1), nn.Conv2d(64, 128, kernel_size=3, stride=2, padding=1), nn.BatchNorm2d(128), nn.ReLU(inplace=True), nn.AdaptiveAvgPool2d((1, 1)), ) self.inter_frame_attn = nn.MultiheadAttention( embed_dim=feat_dim, num_heads=num_heads, batch_first=True ) self.diff_encoder = nn.Sequential( nn.Linear(feat_dim, feat_dim // 2), nn.ReLU(inplace=True), nn.Linear(feat_dim // 2, feat_dim), ) self.gaze_regressor = nn.Sequential( nn.Linear(feat_dim * 2, feat_dim), nn.ReLU(inplace=True), nn.Dropout(0.3), nn.Linear(feat_dim, 3) ) self.norm = nn.LayerNorm(feat_dim) def forward(self, video_clip): """ 前向传播 Args: video_clip: 视频片段, shape=(B, T, C, H, W) B=batch, T=帧数, C=3, H/W=分辨率 Returns: gaze_pred: 视线方向预测, shape=(B, T, 3) attn_weights: 注意力权重, shape=(B, T, T) """ B, T, C, H, W = video_clip.shape frames_flat = video_clip.view(B * T, C, H, W) frame_feats = self.frame_encoder(frames_flat) frame_feats = frame_feats.view(B, T, -1) if frame_feats.size(-1) != self.feat_dim: frame_feats = F.pad( frame_feats, (0, self.feat_dim - frame_feats.size(-1)) ) frame_diffs = torch.zeros_like(frame_feats) frame_diffs[:, 1:] = frame_feats[:, 1:] - frame_feats[:, :-1] diff_embeds = self.diff_encoder(frame_diffs) fused = self.norm(frame_feats + diff_embeds) attn_out, attn_weights = self.inter_frame_attn( fused, fused, fused ) attn_out = self.norm(attn_out + frame_feats) gaze_input = torch.cat([attn_out, diff_embeds], dim=-1) gaze_pred = self.gaze_regressor(gaze_input) gaze_pred = F.normalize(gaze_pred, p=2, dim=-1) return gaze_pred, attn_weights
class FIFALoss(nn.Module): """FIFA 训练损失函数""" def __init__(self, alpha=1.0, beta=0.3): super().__init__() self.alpha = alpha self.beta = beta def forward(self, gaze_pred, gaze_gt, attn_weights): """ Args: gaze_pred: 预测视线, (B, T, 3) gaze_gt: 真实视线, (B, T, 3) attn_weights: 注意力权重, (B, T, T) """ cos_sim = F.cosine_similarity(gaze_pred, gaze_gt, dim=-1) angle_loss = (1 - cos_sim).mean() temporal_diff = gaze_pred[:, 1:] - gaze_pred[:, :-1] temporal_loss = torch.norm(temporal_diff, p=2, dim=-1).mean() attn_entropy = -(attn_weights * torch.log(attn_weights + 1e-8)).sum(-1).mean() total_loss = ( self.alpha * angle_loss + self.beta * temporal_loss - 0.01 * attn_entropy ) return total_loss, { 'angle_loss': angle_loss.item(), 'temporal_loss': temporal_loss.item(), 'attn_entropy': attn_entropy.item() }
if __name__ == "__main__": batch_size = 4 num_frames = 5 model = FineGrainedInterFrameAttention( feat_dim=512, num_heads=8, window_size=5 ) video = torch.randn(batch_size, num_frames, 3, 224, 224) gaze_pred, attn_weights = model(video) print(f"输入形状: {video.shape}") print(f"视线预测形状: {gaze_pred.shape}") print(f"注意力权重形状: {attn_weights.shape}") print(f"预测视线向量范数: {torch.norm(gaze_pred, p=2, dim=-1)[0]}") criterion = FIFALoss() gaze_gt = F.normalize(torch.randn_like(gaze_pred), p=2, dim=-1) loss, metrics = criterion(gaze_pred, gaze_gt, attn_weights) print(f"\n损失: {loss.item():.4f}") for k, v in metrics.items(): print(f" {k}: {v:.4f}")
|