TransGaze-Object:基于Transformer的驾驶员视线物体直接预测框架

论文信息

核心创新

跳过中间的”视线落点(Point-of-Gaze)估计→物体关联”两步流程,直接从驾驶员面部图像+交通场景图像预测视线关注的物体类别。使用Transformer架构实现跨模态融合。

方法架构

graph LR
    A[驾驶员面部图像] --> B[人脸特征提取]
    C[交通场景图像] --> D[场景特征提取]
    B --> E[Transformer编码器]
    D --> E
    E --> F[跨模态注意力]
    F --> G[视线物体分类]
    G --> H[预测物体类别]

传统方法 vs TransGaze-Object

方法 步骤1 步骤2 步骤3 累积误差
传统3步 视线方向估计 落点投影 物体关联 高(每步误差累积)
传统2步 视线落点估计 物体关联 — 中
TransGaze-Object 直接预测物体 — — 低

代码实现

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
import torch
import torch.nn as nn
import torch.nn.functional as F

class TransGazeObject(nn.Module):
"""
TransGaze-Object: 驾驶员视线物体直接预测

输入: 驾驶员面部 + 交通场景
输出: 视线关注的物体类别

架构:
1. 人脸编码器: 提取头部姿态和眼睛方向特征
2. 场景编码器: 提取交通场景物体特征
3. Transformer跨模态融合
4. 分类头预测视线物体
"""

def __init__(self, num_objects: int = 15,
feat_dim: int = 256, num_heads: int = 8,
num_layers: int = 4):
super().__init__()

# 人脸编码器 (简化版CNN)
self.face_encoder = nn.Sequential(
nn.Conv2d(3, 32, 3, stride=2, padding=1),
nn.BatchNorm2d(32), nn.ReLU(),
nn.Conv2d(32, 64, 3, stride=2, padding=1),
nn.BatchNorm2d(64), nn.ReLU(),
nn.Conv2d(64, 128, 3, stride=2, padding=1),
nn.BatchNorm2d(128), nn.ReLU(),
nn.AdaptiveAvgPool2d((4, 4)),
nn.Flatten(),
nn.Linear(128 * 16, feat_dim)
)

# 场景编码器 (简化版CNN)
self.scene_encoder = nn.Sequential(
nn.Conv2d(3, 32, 3, stride=2, padding=1),
nn.BatchNorm2d(32), nn.ReLU(),
nn.Conv2d(32, 64, 3, stride=2, padding=1),
nn.BatchNorm2d(64), nn.ReLU(),
nn.Conv2d(64, 128, 3, stride=2, padding=1),
nn.BatchNorm2d(128), nn.ReLU(),
nn.Conv2d(128, 256, 3, stride=2, padding=1),
nn.BatchNorm2d(256), nn.ReLU(),
nn.AdaptiveAvgPool2d((8, 8)),
nn.Flatten(),
nn.Linear(256 * 64, feat_dim * 4)
)

# 场景物体token化
self.object_tokenizer = nn.Linear(feat_dim * 4, 8 * feat_dim)

# 人脸query token
self.face_proj = nn.Linear(feat_dim, feat_dim)

# Transformer
encoder_layer = nn.TransformerEncoderLayer(
d_model=feat_dim, nhead=num_heads,
dim_feedforward=feat_dim * 4,
dropout=0.1, batch_first=True, activation='gelu'
)
self.transformer = nn.TransformerEncoder(encoder_layer, num_layers)

# 分类头
self.classifier = nn.Sequential(
nn.LayerNorm(feat_dim),
nn.Linear(feat_dim, 64),
nn.GELU(),
nn.Dropout(0.1),
nn.Linear(64, num_objects)
)

def forward(self, face_img: torch.Tensor,
scene_img: torch.Tensor) -> torch.Tensor:
"""
Args:
face_img: (B, 3, 224, 224) 驾驶员面部
scene_img: (B, 3, 480, 640) 前方交通场景

Returns:
logits: (B, num_objects) 视线物体类别
"""
B = face_img.shape[0]

# 编码
face_feat = self.face_encoder(face_img) # (B, feat_dim)
scene_feat = self.scene_encoder(scene_img) # (B, feat_dim*4)

# Token化场景
object_tokens = self.object_tokenizer(scene_feat) # (B, 8*feat_dim)
object_tokens = object_tokens.view(B, 8, -1) # (B, 8, feat_dim)

# 人脸作为query token
face_token = self.face_proj(face_feat).unsqueeze(1) # (B, 1, feat_dim)

# 拼接: [face_token, object_tokens]
tokens = torch.cat([face_token, object_tokens], dim=1) # (B, 9, feat_dim)

# Transformer
encoded = self.transformer(tokens) # (B, 9, feat_dim)

# 取face token的输出做分类
face_output = encoded[:, 0, :] # (B, feat_dim)
logits = self.classifier(face_output)

return logits


# ==================== 测试 ====================
if __name__ == "__main__":
print("=" * 60)
print("TransGaze-Object: 视线物体直接预测")
print("=" * 60)

# 视线物体类别(典型驾驶场景)
gaze_objects = [
"road_ahead", "left_mirror", "right_mirror", "rear_mirror",
"center_console", "infotainment", "speedometer",
"steering_wheel", "phone", "passenger", "window_left",
"window_right", "door_panel", "gear_shift", "other"
]

model = TransGazeObject(
num_objects=len(gaze_objects),
feat_dim=256, num_heads=8, num_layers=4
)

params = sum(p.numel() for p in model.parameters())
print(f"\n模型参数: {params:,} ({params/1e6:.2f}M)")

# 模拟输入
B = 4
face = torch.randn(B, 3, 224, 224)
scene = torch.randn(B, 3, 480, 640)

model.eval()
with torch.no_grad():
logits = model(face, scene)

probs = F.softmax(logits, dim=1)
preds = torch.argmax(probs, dim=1)

print(f"\n预测结果:")
for i in range(B):
obj = gaze_objects[preds[i]]
conf = probs[i][preds[i]].item()
print(f" Sample {i}: {obj:<20} (置信度: {conf:.2%})")

print(f"\n{'='*60}")
print("方法对比")
print(f"{'='*60}")

comparison = [
("3-step pipeline", "视线→落点→关联", "高累积误差", "~120ms"),
("2-step pipeline", "视线落点→关联", "中等误差", "~80ms"),
("TransGaze-Object", "直接预测物体", "低误差", "~45ms"),
]

print(f"{'方法':<20} {'流程':<20} {'误差':<15} {'延迟':<10}")
print("-" * 65)
for m, f, e, l in comparison:
print(f"{m:<20} {f:<20} {e:<15} {l:<10}")

IMS应用启示

1. 分心检测直接输出物体类别

Euro NCAP场景 传统方法输出 TransGaze-Object输出
D-02手机使用 “视线偏离道路” “phone”(直接分类)
D-03打字 “视线向下+手部移动” “phone” + “center_console”
D-05视线偏离 “视线偏离>3秒” “window_left”/“passenger”
正常驾驶 “视线在道路” “road_ahead”

2. 开发建议

优先级 建议 输入 输出 价值
🔴 P0 视线物体分类替代角度回归 人脸+前方场景 15类物体 直接对应NCAP场景
🟡 P1 实时场景物体token化 前方摄像头 8个物体token 跨模态注意力
🟢 P2 多帧时序建模 视频序列 时序注意力 减少误报

总结

TransGaze-Object代表了DMS分心检测从”几何视线估计”到”语义物体理解”的范式转变。跳过中间步骤直接预测视线关注的物体,减少累积误差,输出直接对应Euro NCAP分心场景。对IMS:用物体分类替代角度阈值,可直接输出”手机使用”、”中控操作”等可执行的分心标签。


https://dapalm.com/2026/10/08/2026-10-08-007-transgaze-object-transformer-gaze-prediction/
作者
Mars
发布于
2026年10月8日
许可协议