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
| import torch import torch.nn as nn
class EyeCue(nn.Module): """ EyeCue: 基于眼动-视频融合的认知分心检测框架 架构: 1. Video Encoder: 提取场景时空特征 2. Gaze Encoder: 建模时序眼动模式 3. GDSQ Module: 眼动驱动的语义查询,建模交互 """ def __init__(self, config): super().__init__() self.video_encoder = TimeSformer( img_size=224, patch_size=16, num_frames=16, embed_dim=768, depth=12, num_heads=12 ) self.gaze_encoder = nn.TransformerEncoder( nn.TransformerEncoderLayer( d_model=256, nhead=8, dim_feedforward=1024, dropout=0.1 ), num_layers=6 ) self.gdsq = GazeDrivenSemanticQuery( gaze_dim=256, video_dim=768, hidden_dim=512 ) self.classifier = nn.Sequential( nn.Linear(768 + 256 + 512, 512), nn.ReLU(), nn.Dropout(0.5), nn.Linear(512, 2) ) def forward(self, video, gaze_points): """ Args: video: (B, T, C, H, W) 视频序列 gaze_points: (B, T, 2) 注视点坐标序列 Returns: logits: (B, 2) 分类结果 """ video_features = self.video_encoder(video) gaze_embed = self._embed_gaze(gaze_points) gaze_features = self.gaze_encoder(gaze_embed) interaction_features = self.gdsq(gaze_features, video_features) combined = torch.cat([ video_features.mean(dim=1), gaze_features.mean(dim=1), interaction_features ], dim=-1) return self.classifier(combined) def _embed_gaze(self, gaze_points): """将2D注视点嵌入到高维空间""" B, T, _ = gaze_points.shape pos_embed = self._positional_encoding(T).unsqueeze(0).expand(B, -1, -1) gaze_embed = self.gaze_proj(gaze_points) + pos_embed return gaze_embed
|