因果表征学习:视线估计跨域泛化突破

研究背景

视线估计的域差距问题是实际部署的核心挑战。在特定数据集上训练的模型,应用到新环境时性能严重下降。

域差距来源

差距类型 具体表现 影响
主体外观 不同人种、眼镜、妆容 🔴 高
图像质量 分辨率、压缩比、传感器差异 🟡 中
拍摄角度 摄像头安装位置差异 🔴 高
光照条件 日间/夜间/逆光场景 🔴 高
设备差异 不同的摄像头型号 🟡 中

传统解决方案的局限

域适应(Domain Adaptation) 需要:

  • 目标域样本(实际部署时难以获取)
  • 额外的优化过程
  • 可能影响用户体验

本文创新:域泛化(Domain Generalization)

  • 不需要目标域数据
  • 零样本跨域推理
  • 更适合实际部署场景

核心创新:因果表征学习

因果机制的一般原则

论文首次将因果推理引入视线估计领域,基于因果机制的四大原则:

graph TD
    A[因果机制原则] --> B[共同原因原则]
    A --> C[稳定性原则]
    A --> D[模块化原则]
    A --> E[因果异质性原则]
    
    B --> B1[Y和X相关则存在共同原因C]
    C --> C1[因果关系在不同环境下稳定]
    D --> D1[改变一个机制不影响其他机制]
    E --> E1[不同因果因子重要性不同]

CauGE框架

Causal Representation-Based Domain Generalization on Gaze Estimation

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

class CauGE(nn.Module):
"""
因果表征学习框架

核心思想:
1. 分离因果因子和非因果因子
2. 只保留因果因子用于视线估计
3. 通过对抗训练提取域不变特征
"""

def __init__(self, backbone='ResNet18', feature_dim=256):
super().__init__()

# 特征提取器
self.encoder = ResNetEncoder(backbone, feature_dim)

# 因果因子提取器
self.causal_projector = nn.Sequential(
nn.Linear(feature_dim, feature_dim),
nn.ReLU(),
nn.Linear(feature_dim, feature_dim)
)

# 非因果因子提取器(用于对抗训练)
self.non_causal_projector = nn.Sequential(
nn.Linear(feature_dim, feature_dim),
nn.ReLU(),
nn.Linear(feature_dim, feature_dim)
)

# 视线回归器
self.gaze_regressor = nn.Linear(feature_dim, 2) # pitch, yaw

# 域判别器(对抗训练)
self.domain_discriminator = nn.Sequential(
nn.Linear(feature_dim, 128),
nn.ReLU(),
nn.Linear(128, num_domains)
)

# 注意力层(突出视线相关特征)
self.attention = nn.Sequential(
nn.Linear(feature_dim, feature_dim),
nn.Sigmoid()
)

def forward(self, x, domain_label=None):
"""
Args:
x: (B, C, H, W) 人脸图像
domain_label: 域标签(用于对抗训练)

Returns:
gaze: (B, 2) 视线方向 (pitch, yaw)
domain_pred: 域预测(用于对抗损失)
"""
# 特征提取
features = self.encoder(x)

# 分离因果和非因果因子
causal_features = self.causal_projector(features)
non_causal_features = self.non_causal_projector(features)

# 注意力加权
attention_weights = self.attention(causal_features)
attended_features = causal_features * attention_weights

# 视线回归
gaze = self.gaze_regressor(attended_features)

# 域判别(用于对抗训练)
domain_pred = self.domain_discriminator(non_causal_features)

return gaze, domain_pred, causal_features, non_causal_features


class CauGELoss(nn.Module):
"""
CauGE损失函数

组成:
1. 视线回归损失
2. 对抗损失(域判别)
3. 因果约束损失(独立性)
4. 稳定性损失(减小方差)
"""

def __init__(self, lambda_adv=1.0, lambda_causal=0.5, lambda_stable=0.1):
super().__init__()
self.lambda_adv = lambda_adv
self.lambda_causal = lambda_causal
self.lambda_stable = lambda_stable

def forward(self, gaze_pred, gaze_gt, domain_pred, domain_gt,
causal_features, non_causal_features):

# 1. 视线回归损失
gaze_loss = F.mse_loss(gaze_pred, gaze_gt)

# 2. 对抗损失(梯度反转)
# 最大化域判别器损失,使因果特征与域无关
domain_loss = F.cross_entropy(domain_pred, domain_gt)
adversarial_loss = -domain_loss # 梯度反转

# 3. 因果约束损失:因果和非因果特征应独立
# 协方差矩阵应为零(独立性)
cov = torch.mean(
causal_features * non_causal_features, dim=0
)
causal_loss = torch.norm(cov, p=2)

# 4. 稳定性损失:同一主体的特征方差应小
# 使用对比学习思想
stable_loss = self.stability_loss(causal_features)

# 总损失
total_loss = (
gaze_loss +
self.lambda_adv * adversarial_loss +
self.lambda_causal * causal_loss +
self.lambda_stable * stable_loss
)

return total_loss, {
'gaze_loss': gaze_loss.item(),
'adversarial_loss': adversarial_loss.item(),
'causal_loss': causal_loss.item(),
'stable_loss': stable_loss.item()
}

def stability_loss(self, features):
"""稳定性损失:减小同类特征方差"""
# L2正则化
return torch.norm(features, p=2, dim=1).mean()


# 梯度反转层(GRL)
class GradientReversalLayer(torch.autograd.Function):
"""
梯度反转层
前向传播:恒等变换
反向传播:梯度取反
"""

@staticmethod
def forward(ctx, x, lambda_):
ctx.lambda_ = lambda_
return x.clone()

@staticmethod
def backward(ctx, grad_output):
return -ctx.lambda_ * grad_output, None

因果机制四大原则

1. 共同原因原则(Common Cause Principle)

如果Y(视线)和X(图像特征)相关,则存在共同原因C(真实视线方向)。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
def common_cause_principle(features, gaze):
"""
验证共同原因原则

思想:视线相关的特征应该与视线方向高度相关
"""
# 计算特征与视线的相关性
correlation = torch.corrcoef(
torch.cat([features, gaze], dim=1)
)[:features.shape[1], features.shape[1]:]

# 高相关性特征应保留(因果因子)
# 低相关性特征应去除(非因果因子)

return correlation

2. 稳定性原则(Stability)

因果关系在不同环境下应保持稳定。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
def stability_principle(features_source, features_target):
"""
验证稳定性原则

思想:因果特征在不同域上的分布应相似
"""
# 计算源域和目标域特征的分布差异
source_mean = features_source.mean(dim=0)
target_mean = features_target.mean(dim=0)

source_std = features_source.std(dim=0)
target_std = features_target.std(dim=0)

# 稳定的因果特征应该分布差异小
distribution_shift = (
torch.norm(source_mean - target_mean, p=2) +
torch.norm(source_std - target_std, p=2)
)

return distribution_shift

3. 模块化原则(Modularity)

改变一个因果机制不应直接影响其他机制。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
def modularity_principle(causal_features, non_causal_features):
"""
验证模块化原则

思想:因果和非因果特征应相互独立
"""
# 计算协方差矩阵
cov_matrix = torch.cov(
torch.cat([causal_features, non_causal_features], dim=1).T
)

# 因果和非因果之间的协方差应接近零
cross_cov = cov_matrix[
:causal_features.shape[1],
causal_features.shape[1]:
]

independence_score = torch.norm(cross_cov, p='fro')

return independence_score

4. 因果异质性原则(Causal Heterogeneity)

不同因果因子的重要性不同。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
def causal_heterogeneity_principle(features, gaze):
"""
验证因果异质性原则

思想:眼周区域特征的重要性远高于其他区域
"""
# 计算每个特征的重要性(梯度)
features.requires_grad = True
gaze_pred = model(features)

# 计算视线对每个特征的梯度
importance = torch.autograd.grad(
gaze_pred.sum(), features
)[0].abs().mean(dim=0)

# 眼周区域特征应有更高的重要性权重
return importance

跨域泛化实验结果

实验设置

数据集:

  • ETH-XGaze(训练)
  • MPIIFaceGaze(测试)
  • Gaze360(测试)
  • RT-GENE(测试)

性能对比

方法 ETH→MPII ETH→Gaze360 ETH→RT-GENE 平均
Baseline (vanilla) 7.2° 12.5° 9.8° 9.8°
Domain Adaptation 6.1° 10.2° 8.5° 8.3°
LatentGaze 5.8° 9.9° 8.1° 7.9°
3DGazeNet 5.5° 9.7° 7.9° 7.7°
CauGE (本文) 4.9° 8.8° 7.2° 7.0°

消融实验

组件 ETH→MPII ETH→Gaze360 说明
完整模型 4.9° 8.8° 全部组件
去除对抗训练 5.6° 9.5° +0.7° 性能下降
去除因果约束 5.3° 9.2° +0.4° 性能下降
去除注意力层 5.1° 9.0° +0.2° 性能下降

IMS开发启示

1. 跨车型部署策略

传统方案: 每个车型单独训练模型

CauGE方案: 一次训练,跨车型泛化

1
2
3
4
5
6
7
8
# 部署配置
deployment_config = {
'training_domains': [' sedan_A', 'SUV_B', 'truck_C'],
'target_domains': ['sedan_D', 'SUV_E', 'sports_F'],
'approach': 'domain_generalization',
'target_accuracy': '< 5° error',
'deployment_cost': 'low', # 无需目标域数据
}

2. 数据采集优化

基于因果异质性原则,数据采集应重点关注:

因果因子 重要性 采集策略
眼周区域特征 🔴 极高 高分辨率采集,覆盖眼镜/墨镜
头部姿态 🟡 中等 多角度拍摄,覆盖不同坐姿
光照条件 🟡 中等 覆盖日间/夜间/逆光场景
服装/发型 🟢 低 可忽略,非因果因子

3. 域适应验证测试

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
### DG-01 跨域视线估计测试

**前置条件:**
- 模型在ETH-XGaze数据集训练
- 目标车型为新车型的DMS摄像头数据

**测试步骤:**
1. 在目标车型采集10名驾驶员数据(仅验证,不训练)
2. 计算跨域视线估计误差
3. 对比有/无域适应的性能差异

**判定条件:**
| 指标 | 通过条件 | 失败条件 |
|------|---------|---------|
| 角度误差 | ≤ 5° | > 5° |
| 相比训练域性能下降 | ≤ 20% | > 20% |
| 推理时延 | ≤ 30ms | > 30ms |

**预期输出:**

[跨域测试] 源域性能: 4.2°
[跨域测试] 目标域性能: 4.8°
[跨域测试] 性能下降: 14.3% ✅

1


参考文献

  • 论文标题:Causal Representation-Based Domain Generalization on Gaze Estimation
  • 发表会议:arXiv 2024
  • 论文链接:https://arxiv.org/html/2408.16964v1
  • 代码仓库:待开源

总结

CauGE框架首次将因果推理引入视线估计领域,实现:

  1. 跨域泛化SOTA:无需目标域数据即可在新环境保持高精度
  2. 因果机制指导:基于共同原因、稳定性、模块化、因果异质性四大原则
  3. 对抗训练分离:自动分离因果因子和非因果因子
  4. 实际部署友好:零样本跨域,无需目标域适配

IMS开发优先级: 🔴 高优先级(解决DMS跨车型部署的核心痛点)


因果表征学习:视线估计跨域泛化突破
https://dapalm.com/2026/07/31/2026-07-31-causal-domain-generalization-gaze-estimation/
作者
Mars
发布于
2026年7月31日
许可协议