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
| import torch import torch.nn as nn import torch.nn.functional as F
class GazeSymCAT(nn.Module): """ GazeSymCAT: 对称跨注意力视线估计 核心创新: 1. 双眼对称注意力 2. 跨眼特征交互 3. 极端姿态鲁棒 """ def __init__(self, embed_dim=256, num_heads=8): super().__init__() self.eye_encoder = EyeEncoder() self.cross_attention = CrossAttention(embed_dim, num_heads) self.symmetric_attention = SymmetricAttention(embed_dim, num_heads) self.head_pose_encoder = HeadPoseEncoder() self.gaze_head = nn.Linear(embed_dim * 2 + 128, 3) def forward(self, left_eye, right_eye, head_pose): """ Args: left_eye: (B, 3, 64, 32) 左眼图像 right_eye: (B, 3, 64, 32) 右眼图像 head_pose: (B, 3) 头部姿态(pitch, yaw, roll) Returns: gaze_vector: (B, 3) 视线方向 """ left_feat = self.eye_encoder(left_eye) right_feat = self.eye_encoder(right_eye) left_cross = self.cross_attention(left_feat, right_feat) right_cross = self.cross_attention(right_feat, left_feat) fused = self.symmetric_attention(left_cross, right_cross) pose_feat = self.head_pose_encoder(head_pose) combined = torch.cat([fused, pose_feat], dim=-1) gaze = self.gaze_head(combined) return F.normalize(gaze, dim=-1)
class CrossAttention(nn.Module): """ 跨眼注意力模块 一只眼的特征关注另一只眼 """ def __init__(self, embed_dim, num_heads): super().__init__() self.attention = nn.MultiheadAttention(embed_dim, num_heads, batch_first=True) def forward(self, query, reference): """ Args: query: (B, D) 查询眼特征 reference: (B, D) 参考眼特征 Returns: attended: (B, D) 注意力输出 """ query = query.unsqueeze(1) reference = reference.unsqueeze(1) attended, _ = self.attention(query, reference, reference) return attended.squeeze(1)
class SymmetricAttention(nn.Module): """ 对称注意力模块 左右眼特征对称融合 """ def __init__(self, embed_dim, num_heads): super().__init__() self.self_attn = nn.MultiheadAttention(embed_dim, num_heads, batch_first=True) self.norm = nn.LayerNorm(embed_dim) def forward(self, left_feat, right_feat): """ Args: left_feat: (B, D) right_feat: (B, D) Returns: fused: (B, D) """ sequence = torch.stack([left_feat, right_feat], dim=1) attended, _ = self.self_attn(sequence, sequence, sequence) sequence = self.norm(sequence + attended) fused = sequence.mean(dim=1) return fused
|