TransGaze-Object:Transformer 驾驶员注视物体预测框架与 UD-FSG 数据集

TransGaze-Object:Transformer 驾驶员注视物体预测框架与 UD-FSG 数据集

论文信息

项目 内容
标题 TransGaze-Object: Transformer Based Driver Gaze Object Prediction Framework in Real Driving
arXiv 2609.10139
日期 2026-09-09
页数 32页, 17图
领域 cs.CV

核心创新

首次提出端到端驾驶员注视物体预测框架,跳过中间的注视点(PoG)估计,直接从面部图像预测驾驶员注视的交通物体:

传统方法 TransGaze-Object
面部→PoG→关联物体 面部+场景→直接预测物体
51% 精度 60% 精度
23.21% 背景混淆 11.68% (-49.7%)

方法详解

架构

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

class TransGazeObject(nn.Module):
"""
TransGaze-Object: 端到端驾驶员注视物体预测

输入: 驾驶员面部图像 + 交通场景图像 + 场景物体边界框
输出: 注视物体分类(车辆/行人/信号灯/背景等)

架构:
1. 面部编码器: 提取面部+虹膜特征
2. 场景物体编码器: 提取交通物体空间特征
3. Cross-Attention: 面部特征与物体特征交互
4. 分类头: 预测注视物体
"""

def __init__(self, n_gaze_objects=6, face_feat_dim=512,
object_feat_dim=256, n_heads=8):
super().__init__()

# 1. 面部编码器(ResNet-18 + 虹膜分支)
self.face_encoder = nn.Sequential(
nn.Conv2d(3, 64, 7, 2, 3),
nn.BatchNorm2d(64), nn.ReLU(),
nn.MaxPool2d(3, 2, 1),
*self._make_res_block(64, 128, 2),
*self._make_res_block(128, 256, 2),
nn.AdaptiveAvgPool2d(1)
)

# 虹膜特征分支
self.iris_encoder = nn.Sequential(
nn.Conv2d(1, 32, 3, 1, 1),
nn.BatchNorm2d(32), nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, 3, 1, 1),
nn.BatchNorm2d(64), nn.ReLU(),
nn.AdaptiveAvgPool2d(1)
)

# 融合面部+虹膜
self.face_proj = nn.Linear(256 + 64, face_feat_dim)

# 2. 场景物体编码器
self.object_encoder = nn.Sequential(
nn.Linear(5, 128), # [x1,y1,x2,y2,area]
nn.ReLU(),
nn.Linear(128, object_feat_dim)
)

# 3. Cross-Attention(面部查询物体)
self.cross_attention = nn.MultiheadAttention(
embed_dim=face_feat_dim,
num_heads=n_heads,
kdim=object_feat_dim,
vdim=object_feat_dim,
batch_first=True
)
self.attn_norm = nn.LayerNorm(face_feat_dim)

# 4. 分类头
self.classifier = nn.Sequential(
nn.Linear(face_feat_dim, 128),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(128, n_gaze_objects)
)

def _make_res_block(self, in_c, out_c, n_blocks):
blocks = []
for _ in range(n_blocks):
blocks.extend([
nn.Conv2d(in_c, out_c, 3, 1, 1),
nn.BatchNorm2c(out_c), nn.ReLU(),
nn.Conv2d(out_c, out_c, 3, 1, 1),
nn.BatchNorm2d(out_c), nn.ReLU(),
])
in_c = out_c
return blocks

def forward(self, face_img, iris_img, object_boxes):
"""
Args:
face_img: (B, 3, 224, 224) 驾驶员面部
iris_img: (B, 1, 64, 64) 虹膜区域
object_boxes: (B, N, 5) 场景物体 [x1,y1,x2,y2,area]
Returns:
logits: (B, N, n_gaze_objects) 每个物体的注视概率
"""
B, N, _ = object_boxes.shape

# 面部特征
face_feat = self.face_encoder(face_img).flatten(1) # (B, 256)
iris_feat = self.iris_encoder(iris_img).flatten(1) # (B, 64)
face_fused = self.face_proj(
torch.cat([face_feat, iris_feat], dim=-1)
) # (B, 512)

# 物体特征
obj_feat = self.object_encoder(object_boxes) # (B, N, 256)

# Cross-Attention: 面部查询物体
face_query = face_fused.unsqueeze(1) # (B, 1, 512)
attn_out, attn_weights = self.cross_attention(
face_query, obj_feat, obj_feat
) # (B, 1, 512)
attn_out = self.attn_norm(face_fused.unsqueeze(1) + attn_out)

# 分类
logits = self.classifier(attn_out.squeeze(1)) # (B, n_objects)

# 对每个物体也计算注视概率
all_obj_logits = []
for i in range(N):
obj_attn = attn_weights[:, :, i:i+1] # (B, 1, 1)
obj_feat_i = obj_feat[:, i] * obj_attn.squeeze(-1)
obj_logits = self.classifier(
face_fused + obj_feat_i
)
all_obj_logits.append(obj_logits)

return torch.stack(all_obj_logits, dim=1) # (B, N, n_objects)


# UD-FSG 数据集
UD_FSG_DATASET = {
'name': 'Urban Driving-Face Scene Gaze',
'samples': '~10K synchronized pairs',
'modalities': ['driver_face', 'traffic_scene', 'object_boxes', 'gaze_2d', 'gaze_object'],
'gaze_objects': [
'vehicle', # 车辆
'pedestrian', # 行人
'traffic_signal', # 信号灯
'road', # 道路
'background', # 背景
'unknown', # 未知
],
'collection': 'real urban driving'
}

if __name__ == "__main__":
model = TransGazeObject(n_gaze_objects=6)

# 模拟输入
face = torch.randn(4, 3, 224, 224)
iris = torch.randn(4, 1, 64, 64)
boxes = torch.randn(4, 10, 5) # 10个物体

logits = model(face, iris, boxes)
print(f"面部: {face.shape}")
print(f"虹膜: {iris.shape}")
print(f"物体: {boxes.shape}")
print(f"输出: {logits.shape} (4×10×6)")

UD-FSG 数据集

项目 规格
同步对 驾驶员面部 + 交通场景
标注 场景物体边界框 + 注视坐标 + 注视物体
场景 真实城市驾驶
物体类别 6类(车辆/行人/信号灯/道路/背景/未知)

实验结果

方法 精度 背景混淆率 说明
PoG→物体关联 51% 23.21% 两阶段
TransGaze-Object 60% 11.68% 端到端
提升 +9% -49.7% 显著

IMS 应用启示

1. DMS 注意力监测升级

当前 DMS TransGaze-Object 增强
注视区域(6-9区) 注视物体(语义级)
“看中控” “看手机屏” vs “看导航”
“看右” “看右侧车辆” vs “看路边”
分心判断 基于物体语义判断

2. 部署方案

组件 型号 参数
面部摄像头 OV2311 IR 2MP, 30fps
前向摄像头 已有 ADAS 复用
处理器 QCS8255 26 TOPS
模型 TransGaze-Object INT8 ~15MB

3. 与 ADAS 协同

场景 DMS 输出 ADAS 动作
注视右侧车辆 gaze=vehicle 降低AEB灵敏度
注视信号灯 gaze=signal 绿灯启动延迟
注视行人 gaze=pedestrian 提前刹车准备
注视背景 gaze=background 分心警告

总结

TransGaze-Object 将 DMS 从”注视区域”升级到”注视物体”语义级:

  1. 端到端 > 两阶段:跳过 PoG,直接预测物体
  2. Cross-Attention 是关键:面部特征查询场景物体
  3. 虹膜特征增强:虹膜加权眼部特征提升精度
  4. UD-FSG 数据集开源:真实驾驶标注数据
  5. IMS 可落地:复用已有 DMS+ADAS 摄像头

TransGaze-Object:Transformer 驾驶员注视物体预测框架与 UD-FSG 数据集
https://dapalm.com/2026/09/11/2026-09-11-transgaze-object-transformer-driver-gaze-prediction-ud-fsg-ims/
作者
Mars
发布于
2026年9月11日
许可协议