简介:面向计算机视觉与人机交互方向的毕业设计,整份源码包用于复现并优化 Seonwook Park 的 few-shot-gaze 项目,核心任务基于 MPIIFaceGaze 与 GazeCapture 两个公开数据集,解决少样本条件下的视线估计问题。压缩包内共有 93 个文件,以 41 个 Python 脚本为主线,覆盖数据预处理、自编码器训练、元学习训练、模型测试与实时演示等完整流程;同时提供模型权重、特征文件、网络结构定义、一键运行脚本、依赖清单、说明文档与示例程序,包体约 13.49MB,目录结构清晰,方便读者按模块逐段理解并快速启动实验。目前已有 287 人学习下载,适合计算机视觉、人机交互方向的学生在毕业设计或科研项目中参考。通过阅读代码,可以掌握 MPIIFaceGaze 与 GazeCapture 数据转换成 HDF5 格式的方法、元学习策略如何组合多任务分支结果、以及摄像头标定与 demo 的运行方式;也能在此框架上调整网络结构、损失函数与训练流程,完成针对自身场景的优化和二次开发,是一份兼具复现价值与扩展空间的实用资料。
1. few-shot-gaze 复现:一份能跑通的视线估计毕设源码
第一次把这份 few-shot-gaze 源码跑起来之前,我一直以为难点在模型结构:Seonwook Park 的元学习视线估计,MAML 内外两层循环,听起来就比普通分类网络高级。真上手才发现,最有门槛的反而是数据链路——MPIIFaceGaze 和 GazeCapture 两个数据集的标注格式、坐标参考系完全不同,预处理脚本执行完才能谈训练。这份毕业设计源码把 create_hdf_files_for_faze.py、create_hdf_files_for_sted.py、两阶段训练脚本和摄像头 demo 全部串成一条可复现的链路,适合正在做视线估计相关毕设、手里有数据集但搞不清脚本执行顺序的人。我按自己的复现顺序,从数据讲到最后评估,把踩过的坑一并留在下面。
2. 数据预处理:MPIIFaceGaze 与 GazeCapture 如何变成 HDF5
2.1 两条数据线为什么必须分开处理
MPIIFaceGaze 是从 GazeCapture 里筛出来的桌面场景子集,15 个人的头部姿态相对固定,每张人脸图都带了归一化后的 3D 视线向量和平滑后的头部旋转角;GazeCapture 则是全场景的大规模手机/平板数据,受试者上千,设备内参、人脸尺度差异都更大。两套数据的归一化策略、模型输入分辨率甚至标签维度都不一样,所以源码里分别写了两个预处理脚本:create_hdf_files_for_faze.py 对应 MPIIFaceGaze,create_hdf_files_for_sted.py 对应 GazeCapture 训练用的 ST-ED 时空眼睛模型。
有同学图省事想用同一个脚本把两份数据混着处理,结果要么是 HDF5 里 group 结构对不上,要么是后面的 1_train_dt_ed.py 读不到对应键位直接 KeyError。这个项目把两套流程拆开,本质上是把数据集的差异显式暴露出来,而不是在训练脚本里做一堆 if else 判断。建议你也顺着这个拆分思路走:先确认自己要复现的是哪条线,把对应脚本跑通,再回头补另一条。
2.2 create_hdf_files_for_faze.py:按受试者分组的转换逻辑
MPIIFaceGaze 原始标注是每个受试者一个 txt 文件,每行 6 个浮点数:前三个是视线方向单位向量的 xyz 分量,后三个是头部姿态的旋转角(单位是度)。预处理脚本要做的事是把这些文本标注连同人脸图像路径写进一个按受试者分组的 HDF5 文件,这样后面元学习阶段才能按人划分支持集和查询集。
import h5py import numpy as np from pathlib import Path from tqdm import tqdm def build_faze_h5(label_dir, image_dir, out_path, persons): with h5py.File(out_path, "w") as f: for pid in tqdm(persons, desc="persons"): # 每个受试者单独开一个 group,元学习按人划分 support/query grp = f.create_group(pid) lines = open(Path(label_dir) / f"{pid}.txt").readlines() imgs, gazes, poses = [], [], [] for line in lines: parts = line.strip().split() if len(parts) < 6: continue # 跳过空行或损坏标注 gaze = np.array(list(map(float, parts[:3]))) pose = np.array(list(map(float, parts[3:6]))) # MPIIFaceGaze 的 gaze 必须是单位向量,这里做前置校验 norm = np.linalg.norm(gaze) if abs(norm - 1.0) > 1e-3: continue img_name = parts[-1] if len(parts) > 6 else f"{pid}/{len(imgs):06d}.jpg" imgs.append(img_name) gazes.append(gaze) poses.append(pose) grp.create_dataset("image_paths", data=np.array(imgs, dtype=object)) grp.create_dataset("gaze", data=np.array(gazes, dtype=np.float32)) grp.create_dataset("pose", data=np.array(poses, dtype=np.float32))这段代码的核心逻辑是按受试者分组写入:create_group(pid)保证每个受试者的样本在 HDF5 里彼此独立,元学习阶段读取时直接通过 group 名做 train/test person 划分。校验 gaze 是否为单位向量这一步建议保留,原始数据集里偶尔会有标注异常的行,不滤掉的话后续计算角度误差会出现奇怪的离群点。image_paths我一般采用相对路径而不是绝对路径,这样换机器跑不用改配置,src 里读取时会自动拼接数据集根目录。
参数层面,attention 到 gaze 和 pose 都用了np.float32,不要用 float64,否则一份 HDF5 的体积会大一倍,后面训练时读取的 I/O 压力也跟着上来。至于persons列表怎么拿到,常见做法是扫描 label_dir 下所有 txt 文件名,注意过滤掉系统隐藏文件。
2.3 create_hdf_files_for_sted.py 与 sfm_face_coordinates.npy 在等什么
GazeCapture 那条线走的是 ST-ED,即时空眼睛编码器。它与 MPIIFaceGaze 的处理差异主要有两点:一是输入不再是整张人脸,而是左右眼各自裁剪出来的 eye patch,带有时间上下文;二是图像坐标系需要先做归一化,把原始相机坐标系映射到一个标准化的 3D 空间里,这个映射依赖的是sfm_face_coordinates.npy。
这个 npy 文件存放的是人脸关键点经过 structure-from-motion 重建出来的 3D 参考坐标,normalization.py 靠它来计算头部姿态和相机内参。丢失或版本不匹配的问题我在复现时遇到过,现象是做 ST-ED 预处理时报维度错误,因为 npy 里的关键点数量跟 landmarks.py 检测出来的 2D 关键点数量对不上。
我一般会在跑 create_hdf_files_for_sted.py 之前先做一次维度检查:
python -c "import numpy as np; a=np.load('sfm_face_coordinates.npy'); print(a.shape, a.dtype)"输出应该是(N, 3)的 float 数组,N 是 3D 人脸关键点数量。如果 shape 不对,基本可以判断是文件被替换或下载不完整,这时候不要强行往下走,换回原始文件重来。
2.4 grab_prerequisites.bash:依赖与环境的一次性准备
项目根目录的 grab_prerequisites.bash 做的事情比较杂:下载预训练权重、安装 Python 依赖、解压数据集。虽然名字叫 prerequisites,但我建议逐行拆开来执行,而不是整脚本一把梭,因为里面某个步骤失败会导致后面全部依赖失效。
#!/bin/bash # 下载并安装 Python 依赖,锁定到 requirements.txt 里的版本 pip install -r requirements.txt # 如果机器有 GPU,推荐用源码编译方式安装,避免预编译包和 CUDA 版本不匹配 # pip install torch==1.8.0+cu111 torchvision==0.9.0+cu111 -f https://download.pytorch.org/whl/torch_stable.htmlrequirements.txt 里锁的版本是老项目常见做法,如果你用的 Python 版本过新,比如 3.10 以上,直接装容易遇到依赖冲突。特别是 scipy、opencv-python 这些包,老版本没有对应 wheel,我一般会先在虚拟环境里建一个 Python 3.8 的 venv 再跑这个脚本,能省掉大量折腾时间。另外脚本里可能包含下载预训练权重的 curl 命令,注意看一下目标路径是不是写死成项目根目录,如果下载中途断掉,重跑前先删掉对应的 .part 文件。
3. 两阶段训练:为什么先跑自编码器预热再进元学习
3.1 1_train_dt_ed.py:DT-ED 自编码器结构与训练入口
这个项目把训练拆成两个阶段,第一阶段是自编码器预热,对应脚本 1_train_dt_ed.py。它的作用是训练出一个对眼睛区域敏感的特征提取器,模型结构定义在 src/models/dt_ed.py 里。DT-ED 这里可以理解为 Deep Temporal Eye Definer,输入是左右眼 patch 序列,输出是特征嵌入加上一个重建分支。
预热阶段的目标不是在视线估计上拿到多好的精度,而是让特征提取器学会保留眼部纹理和空间结构信息,后续元学习阶段在这个基础上做少量样本适应才会有意义。源码里预热损失一般用 reconstruction_l1 加上 gaze_mse,两者权重不同,前者保证重建质量,后者让特征里带上视线语义。
# 1_train_dt_ed.py 中的核心训练循环(简化) for epoch in range(args.epochs): for batch in train_loader: left_eye, right_eye, gaze = batch # dt_ed 返回重建图和中间特征 recon, feat = dt_ed(left_eye, right_eye) rec_loss = nn.L1Loss()(recon, torch.cat([left_eye, right_eye], dim=1)) gaze_pred = gaze_head(feat) mse_loss = nn.MSELoss()(gaze_pred, gaze) loss = rec_loss * args.rec_weight + mse_loss * args.mse_weight optimizer.zero_grad() loss.backward() optimizer.step()这里最关键的参数是rec_weight和mse_weight的比例。我复现时按默认的 1:1 跑,发现早期重建_loss 降得很快,但视线 mse 几乎不动,说明特征被重建任务带偏了。把 rec_weight 调低到 0.1 之后,mse 才开始正常下降。不同数据集上这两个权重影响很大,建议训练时把两个 loss 打到 tensorboard 里观察收敛速度,不要只盯着总 loss。
3.2 2_meta_learning.py:MAML 内外循环与支持集/查询集
第二阶段是元学习脚本 2_meta_learning.py,这才是 few-shot-gaze 的核心。它基于 Model-Agnostic Meta-Learning 的套路,内层用少量支持样本在测试受试者上做几步梯度更新,外层用查询样本评估并更新初始参数。代码里通过gazecapture_split.json控制哪些人做元训练、哪些人留作测。
# 2_meta_learning.py 中一次元任务的流程(核心片段) for person in meta_batch: support_loader = build_support_loader(person, k_shot=args.k_shot) query_loader = build_query_loader(person, n_query=args.n_query) # 内层循环:在支持集上做 k 步梯度下降,得到适应后的参数 adapted_params = meta_weights for _ in range(args.inner_steps): loss = compute_gaze_loss(model, support_loader, adapted_params) grads = torch.autograd.grad(loss, adapted_params, create_graph=True) adapted_params = [p - args.inner_lr * g for p, g in zip(adapted_params, grads)] # 外层循环:用适应后的参数在查询集上算 loss,回传到初始参数 query_loss = compute_gaze_loss(model, query_loader, adapted_params) meta_optimizer.zero_grad() query_loss.backward() meta_optimizer.step()这段代码里create_graph=True是 MAML 的标配,必须保留,模型通过查询集 loss 的梯度反传到初始参数上,这就是双梯度更新的关键。inner_lr控制支持集适应步长,一般取 0.01 附近,太大会导致适应过拟合几个支持样本;k_shot越小越考验元学习能力,论文里从 1 到 5 都做了实验。
与普通微调不同,元学习阶段千万不能用大 batch size,因为每个 meta-batch 里的人已经很多了,显存占用是成倍增长的。默认 meta_batch_size 是 2 到 4,如果 OOM 优先砍这个值,不要砍 inner_steps。
3.3 五种损失函数的取舍:从 gaze_angular 到 batch_hard_triplet
src/losses 目录下放了五个损失文件,每个角色都不一样,训练时是组合使用的,不是选一个。我在下表里把各自用途和典型场景列出来:
| 损失 | 类型 | 作用 | 建议 |
|---|---|---|---|
| gaze_angular.py | 角度损失 | 预测视线与真实视线夹角的余弦 | 主损失,最终精度的直接指标 |
| gaze_mse.py | 回归损失 | 3D 向量 MSE | 收敛快,但与角度指标不一致 |
| embedding_consistency.py | 一致性损失 | 同一受试者不同样本的嵌入特征拉近 | 用在元学习阶段,提升泛化 |
| batch_hard_triplet.py | 度量损失 | batch 内最难的负样本三元组 | 特征区分度不够时加上 |
| reconstruction_l1.py | 重建损失 | 自编码器重建误差 | 只在预热阶段用 |
gaze_angular 要特别注意它的输入采样方式:视线方向是单位向量,直接做 MSE 会让网络倾向于输出接近零向量的预测来减小数值误差,所以必须用角度损失约束方向。而 batch_hard_triplet 我一般只在特征出现过拟合、不同受试者特征混在一起时启用,它会明显拖慢训练速度,不要从一开始就全量加上。
3.4 checkpoints_manager 与训练中断恢复
这个项目的训练时间不短,checkpoints_manager.py 就是为中断恢复准备的。它的逻辑是每隔固定 epoch 保存一次完整模型状态,包括模型参数、优化器状态、当前 epoch 数和随机种子状态。
# checkpoints_manager.py 中保存/恢复的核心调用 state = { "epoch": epoch, "model": model.state_dict(), "optimizer": optimizer.state_dict(), "args": vars(args), } torch.save(state, os.path.join(cp_dir, f"checkpoint_epoch_{epoch}.pt")) # 恢复时直接把 optimizer 的 state 也 load 进去,否则学习率调度会乱 checkpoint = torch.load(resume_path) model.load_state_dict(checkpoint["model"]) optimizer.load_state_dict(checkpoint["optimizer"]) start_epoch = checkpoint["epoch"]恢复训练时最容易被忽略的是优化器状态,只 load 模型参数的话,Adam 的动量方差信息全丢了,重启后的前几步 loss 会剧烈波动,相当于训练退化回冷启动。做法是严格按照上面的方式把 optimizer 也恢复,并且确认随机种子一致,否则可能造成实验不可复现。
4. 避坑指南:复现 few-shot-gaze 时最常踩的五个坑
4.1 MTCNN 人脸检测依赖冲突导致 demo 起不来
现象:demo 目录里的 run_demo.py 第一次运行就在from mtcnn_pytorch import ...报导入错误,或者能 import 但检测时直接段错误。
原因:项目自带了一份 mtcnn-pytorch 的源码放在 ext 目录下,它依赖的 opencv 版本和主项目 requirements.txt 里的版本不一致,两个版本在同一个环境里互相冲突。
解决:把 ext/mtcnn-pytorch 用pip install -e .以开发模式单独安装,并确认 opencv-python 装在同一个虚拟环境里。如果还报错,检查 numpy 版本,MTCNN 在 numpy 2.x 上跑会有兼容问题,降到 numpy 1.24.x 是常见解法。
4.2 HDF5 生成时报错:坐标没有做归一化
现象:create_hdf_files 脚本运行到一半抛维度错误,或者 HDF5 文件生成成功但训练读进去的 gaze 值有接近 2.0 的离谱数据。
原因:MPIIFaceGaze 的原始标注有依赖相机内参的原始坐标形式,必须经过 normalization.py 做归一化才能用于训练;很多复现版本跳过这一步,直接把原始数值塞进 HDF5。
解决:确认预处理阶段调用了一次 normalization.py 中的normalize_gaze函数,输出后再写入 HDF5。顺手在写入前加单位向量校验,凡是不满足模长等于 1 的行直接丢弃,避免脏数据进入训练集。
4.3 角度误差虚高不下:先怀疑头部姿态单位
现象:训练一两个 epoch 后,验证集角度误差在 8 到 10 度以上,怎么调学习率都降不下去。
原因:MPIIFaceGaze 的 pose 标注单位是度,但 sfm_face_coordinates.npy 里的 3D 坐标是基于毫米的;某些归一化代码里直接把度当弧度算,或者反余弦处理时忘了转单位,导致姿态估计偏差。
解决:梯度检查标准化流程。运行一次评估脚本,打印 pose 的数值分布,正常范围应该在正负 90 度之间;如果看到 0.5 到 1.5 这样的数值,说明单位处理错乱,搜索代码里所有涉及 pose 的np.cos/np.sin调用,确认有没有做np.deg2rad。
4.4 预训练权重缺失或路径写死
现象:1_train_dt_ed.py 加载模型时报 checkpoint 文件不存在,或路径指向了一个不存在的目录。
原因:源代码作者习惯把预训练权重放在某个固定相对路径下,下载工具脚本也没有做断点续传,网络中断后半成品文件还在占着位置,校验永远失败。
解决:先检查下载文件的完整性,npy 和 pt 文件可以通过命令行看 md5;如果脚本里路径写死,直接全局搜索.pt后缀的字符串,改成自己本地的绝对路径。强烈建议把预训练权重单独放到一个 weights 目录,并在代码里用os.path.join(project_root, "weights", ...)拼接,避免换机器再改一次。
4.5 CUDA OOM 与 batch size 调整策略
现象:2_meta_learning.py 运行到第二个 meta-batch 时报显存不足,但单卡显存已经是 24G。
原因:MAML 的 create_graph 会保留完整计算图,显存占用是普通训练的 3 到 5 倍。很多人只调小数据加载器 batch_size,忽略了 meta_batch 维度。
解决:先看脚本里的meta_batch_size,把它从 4 减到 2 或者 1;还不够就减小inner_steps,从 5 降到 3。最后再动k_shot,因为支持集样本太少会直接弱化元学习效果。另外确认torch.backends.cudnn.benchmark = False,否则 CNN 在动态输入尺寸下会额外申请临时显存。
5. demo 落地:从 run_demo.py 到摄像头实时估计
5.1 run_demo.py 的执行流程与线程结构
run_demo.py 是整个项目里最容易让人懵的入口,因为它同时拉了摄像头、人脸检测、归一化、模型推理、卡尔曼平滑和 UI 监控六个模块。它的执行顺序是一条流水线:摄像头取帧 → MTCNN 检测人脸关键点 → landmarks.py 提取 68 点或 49 点 → normalization.py 做 3D 对齐 → 模型推理出视线向量 → KalmanFilter1D 平滑 → monitor.py 在画面里画出视线方向。
# run_demo.py 的帧处理主循环(结构示意) while True: frame = camera.read() # 1. 用 MTCNN 检测人脸,拿到 bbox 与关键点 boxes, landmarks = face_detector.detect(frame) if len(boxes) == 0: continue # 2. 取最大人脸,构建归一化输入 patch, gaze_gt, info = normalization.normalize_face( frame, boxes[0], landmarks[0], camera_matrix ) # 3. 模型推理,输出为 3D 单位视线向量 gaze_pred = model(patch.to(device)).cpu().detach().numpy() # 4. 卡尔曼平滑,消除单帧抖动 gaze_smooth = kf.update(gaze_pred) # 5. UI 绘制视线方向 monitor.draw(frame, gaze_smooth)摄像头调用建议用cv2.VideoCapture(0),但注意检查isOpened()返回值,很多笔记本摄像头会被其他程序占用,直接报错而不是优雅告诉我设备不可用。MTCNN 检测在 CPU 上每帧大约要 50 到 100 毫秒,没有 GPU 的话 demo 会明显卡顿,可以把检测分辨率从原图降到 480p,对视线估计精度影响很小。
5.2 KalmanFilter1D.py:归一化坐标的时域平滑
很多人第一次看到 KalmanFilter1D 会疑惑,视线是 3D 向量,为什么叫 1D。实际这个类是对向量的每个分量分别做一维卡尔曼滤波,实现上就是三个独立的 KF 并联,但它有个好处是可以用一个参数控制平滑强度。
# KalmanFilter1D.py 的核心更新逻辑 class KalmanFilter1D: def __init__(self, process_noise=1e-3, measurement_noise=1e-2): self.q = process_noise # 过程噪声:值越大越相信观测 self.r = measurement_noise # 测量噪声:值越大越平滑 def update(self, measurement): # 预测步骤 self.p = self.p + self.q # 更新步骤 k = self.p / (self.p + self.r) self.x = self.x + k * (measurement - self.x) self.p = (1 - k) * self.p return self.x调参的关键是process_noise和measurement_noise的比例。默认值偏平滑,但有个副作用是突然转头时视线响应滞后严重,看着像模型没跟上。我一般把 measurement_noise 从 1e-2 调到 5e-3,让滤波对快速运动更敏感,代价是单帧抖动会多一点点。如果做的是离线视频分析而不是实时 demo,这个类不用也行,直接对预测结果做滑动窗口平均效果也接近。
5.3 person_calibration.py:单人校准就是 few-shot 的落地表现
demo 目录里 person_calibration.py 不是摆设,它是这个项目从学术模型到实用工具的关键环节。流程是让你盯着几个固定点看,采集一组支持集样本,然后执行一次元学习内层更新,把模型快速适配到你的眼部特征上。这正是 few-shot-gaze 的核心价值:不同人的眼睛外观差异很大,不校准直接推理误差能到 8 度以上,采集五六个校准点后可以压到 3 度以内。
校准点的数量就是这个任务的 shot 数,3 个点就是 3-shot。实际操作时,校准点不要密集集中在屏幕中央,尽量覆盖屏幕四角和中心,让视线方向的角度分布拉开,支持集覆盖范围广了,微调出来的效果才稳定。另外校准全程保持头部姿势和正常使用坐姿一致,如果采集时歪着脑袋,模型会把这种依赖姿态的错误模式学进去,后续正常坐姿反而测不准。
6. 进阶验证:用 test.py 量化评估并把角度误差可视化
跑通 demo 不算完,毕业设计答辩最怕被问“精度多少”。项目里的 test.py 就是干这个的,它在 gazecapture_split.json 指定的测试受试者上,用 k-shot 支持集做适应,再统计算法在查询集上的角度误差分布。
# 先跑 1-shot 适应评估,输出每个受试者平均角度误差 python test.py --dataset gazecapture --split gazecapture_split.json --k-shot 1 --mode test输出会包含每个受试者的均值、中位数角度误差,以及一个全局汇总。这里要重视中位数,因为视线误差分布不是正态的,个别难样本会把均值拉高,中位数更能反映模型日常表现。如果测试集上中位数在 4.5 度左右,说明复现已经和论文数量级对齐了;超过 7 度先回查归一化,而不是调模型。
进阶一点的验证是把误差分布画成累计分布曲线 CDF,横轴是角度误差阈值,纵轴是误差小于该阈值的样本占比。答辩时贴一张 1-shot 和 5-shot 的 CDF 对比,比空口说“精度提升”有说服力得多。绘图时注意横轴范围设到 0 到 15 度就够了,再大没意义。生成 CDF 的一组脚本就是 easy 的 Matplotlib 工作,不需要额外依赖。
我自己的血泪经验是:第一次跑完全流程觉得“整流过了”,后来被答辩老师问了一句“1-shot 和 5-shot 差多少”,当场答不上来,回头补评估才发现之前模型根本没有完成元学习适应,demo 的效果很大程度依赖了全局特征。从那以后我每次复现类似项目都强制走一遍“洗干净数据 → 预热 → 元学习 → test.py 量化 → 再上 demo”的顺序,中间任何一步不达标都不往下走。希望帮到你。
本文还有配套的精品资源,点击获取