Skip to content

Repository files navigation

BadmintonTrackNet

TrackNetV3-based shuttlecock tracking and event analysis

English | 简体中文

基于 TrackNetV3 的羽毛球轨迹检测、轨迹修复、落点预测与关键事件检测项目。

本仓库以原始 TrackNetV3 为基础,保留“TrackNet 逐帧定位 + InpaintNet 轨迹修复”的主流程,并在此之上扩展了面向羽毛球比赛分析的后处理能力:落点预测、落地/击球事件检测、半自动标注工具、BounceNet 事件分类器、批量评估与可视化分析界面。

原论文:TrackNetV3: Enhancing Shuttlecock Tracking with Augmentations and Trajectory Rectification

论文链接:ACM Digital Library

TrackNetV3 network architecture

项目特点

  • TrackNetV3 轨迹检测主干:使用背景估计、热力图预测和时序集成输出羽毛球坐标。
  • InpaintNet 轨迹修复:对遮挡、漏检和不连续轨迹进行补全,使后续事件检测更稳定。
  • 坐标型落点预测tracknet/landing/ 基于 (x, y) 运动判断落点候选,并用 visibility 排除漏检形成的伪稳定零坐标段。
  • 批量落点评估:支持对 train/test 等 split 批量生成 Top-K 落点候选,并在有标签时输出误差统计。
  • 落地/击球事件检测模块tracknet/events/ 使用运动学规则检测 landing、hit、out_of_frame 等关键事件。
  • BounceNet 二阶段分类器:在高召回规则候选基础上,用轨迹特征或视觉 patch 过滤误检、修正事件类型。
  • 半自动标注工具tools/labeling/ 支持 Phase 1 候选预填充、手动修正、批量标注、统计与导出。
  • 误差分析和可视化:保留并扩展 Dash 误差分析界面,同时新增落点预测视频可视化脚本。
Trajectory comparison

与原 TrackNetV3 的关系

原 TrackNetV3 主要解决“每一帧羽毛球在哪里”的问题,本项目进一步面向“比赛片段中关键事件在哪里发生”:

  1. 先用 TrackNet/InpaintNet 得到连续轨迹。
  2. 再用规则或轻量模型从轨迹中识别落地点、击球点、出画面等事件。
  3. 用标注工具快速修正候选结果,形成可训练的事件标签。
  4. 用 BounceNet 或坐标型落点预测器进行批量推理、评估和可视化。

这使仓库从单纯的轨迹检测模型,扩展为一个更完整的羽毛球轨迹分析与事件检测工具链。

环境安装

原始开发环境参考:

Ubuntu 16.04.7 LTS
Python 3.8.7
torch 1.10.0

安装依赖:

git clone https://github.com/ZSHYC/BadmintonTrackNet.git
cd BadmintonTrackNet
pip install -r requirements.txt

开发模式安装(推荐,代码位于 tracknet/ 包中):

pip install -e .

主要依赖包括 torchopencv-pythonnumpypandasdashplotlyPillowtqdmpycocotoolsmatplotlib。命令行入口支持 --device auto|cpu|cuda|mps;默认 auto 会选择可用加速器,否则使用 CPU。

pyproject.tomlrequirements.txt 声明的是兼容版本下限,不是跨平台锁文件。部署或复现实验时,应在目标 Python、CUDA 和操作系统组合上生成自己的约束文件,并保存实际安装版本。

数据准备

本项目沿用 Shuttlecock Trajectory Dataset 的组织方式。建议目录结构如下:

从原始数据集整理时,按官方 TrackNetV3 的约定处理:

  1. ProfessionalAmateur 下的比赛目录合并到 data/train/
  2. Amateur 的三个比赛目录依次重命名为 match24match25match26
  3. 将原始 Test 目录重命名为 data/test/
data/
  train/
    match1/
      csv/
      frame/
      video/
    match2/
    ...
  val/
  test/
    match1/
      csv/
      corrected_csv/
      frame/
      video/

预处理:

python -m tools.preprocessing.preprocess

说明:

  • CSV 字段通常为 Frame, Visibility, X, Y
  • frame/val/ 可由预处理流程生成。
  • 背景中值图会保存到对应 match 或 rally 的 median 文件中。
  • 默认优先使用 match 级 median.npz;如果同一 match 内机位差异明显,可删除该文件,自动退回到各 rally 的 median.npz。官方 README 以 train/match16/median.npz 为例。
  • split 级 .npz 索引缓存记录参数、来源 CSV 大小和修改时间;标签或预测更新后会自动失效并原子重建,无需手工删除。

数据与接口契约

  • Frame 是真实帧标识符,不是 DataFrame 行号;允许从非零帧开始或存在间隔,输出与评估均按该字段对齐。
  • 轨迹坐标在数据读取、事件检测、落点预测、输出和评估阶段保持为原始视频像素,只在模型输入边界显式归一化。
  • Visibility=0 的坐标不会作为有效运动学或落点稳定性证据;缺失点可由 InpaintNet 修复,但不会被静默改成真实观测。
  • 视频宽高和 FPS 随轨迹一起传递,供坐标缩放、边界规则和运动学计算使用。
  • Dataset 返回缓存数组的副本,读取同一样本不会重复归一化或修改缓存内容。

TrackNetV3 轨迹推理

从官方 TrackNetV3 下载预训练权重:

下载 TrackNetV3_ckpts.zip 后解压,并将 TrackNet_best.ptInpaintNet_best.pt 放在 ckpts/

PyTorch checkpoint 可能包含 pickle 数据。仅加载可信来源的权重,并在分发方提供校验值时先核对 SHA-256;不要把未知 .pt 文件传给推理、训练或评估入口。

unzip TrackNetV3_ckpts.zip

视频推理并输出预测 CSV:

python -m tracknet.inference.tracknet \
  --video_file test.mp4 \
  --tracknet_file ckpts/TrackNet_best.pt \
  --inpaintnet_file ckpts/InpaintNet_best.pt \
  --save_dir prediction

同时输出带预测轨迹的视频:

python -m tracknet.inference.tracknet \
  --video_file test.mp4 \
  --tracknet_file ckpts/TrackNet_best.pt \
  --inpaintnet_file ckpts/InpaintNet_best.pt \
  --save_dir prediction \
  --output_video

视频推理默认使用有界窗口流式读取,长视频无需切换另一套算法。旧的 --large_video 参数仍保留兼容;--video_range 只限定背景中值采样区间:

python -m tracknet.inference.tracknet \
  --video_file test.mp4 \
  --tracknet_file ckpts/TrackNet_best.pt \
  --inpaintnet_file ckpts/InpaintNet_best.pt \
  --save_dir prediction \
  --large_video \
  --video_range 324,330

窗口会按真实帧编号集成并在边界重新归一化权重;短于序列长度的视频会用末个真实帧补齐模型输入,但输出只包含真实视频帧。显式指定不可用的 cudamps 会立即报错,不会静默切换设备。

如需同时输出规则事件:

python -m tracknet.inference.events \
  --video_file test.mp4 \
  --tracknet_file ckpts/TrackNet_best.pt \
  --inpaintnet_file ckpts/InpaintNet_best.pt \
  --save_dir prediction-events \
  --output_video

该入口生成 *_ball.csv*_events.json,检测到事件时另生成 *_events.csv--no-detect-bounce 可只运行轨迹推理。事件规则使用视频实际宽高和 FPS,输出事件帧沿用真实 Frame

坐标型落点预测

新增的落点预测器位于:

  • tracknet/landing/predictor.py:单条轨迹的 Top-K 落点候选生成。
  • tracknet/landing/batch.py:按 split/match 批量推理和评估。
  • tracknet/landing/visualize.py:把预测落点画回视频。

落点预测统一通过 tracknet.landing 包使用。

核心策略:

  • visibility 仅用于排除 TrackNet 丢失形成的零坐标伪稳定窗口,不作为落地事件本身的判据。
  • 使用稳定窗口、低速度、Y 坐标稳定性和转向信号筛选候选。
  • 输出 Top-K 候选,按转向、平均速度和 Y 方向波动排序。
  • 有标注时计算帧误差、坐标误差和命中率;有确认标签但没有预测的样本仍保留在命中率分母中。

训练集评估并保存预测:

python -m tracknet.landing.batch --split train --evaluate --save_preds

测试集批量预测:

python -m tracknet.landing.batch --split test --save_preds

常用参数:

python -m tracknet.landing.batch \
  --split test \
  --window 5 \
  --dy 10 \
  --v_th 5 \
  --top_k 5 \
  --save_preds

输出示例:

data/<split>/<match>/pred_landing/
  *_pred.json
  metrics.json   # evaluate 模式下生成
  detail.json    # evaluate 模式下生成

预测结果可视化:

python -m tracknet.landing.visualize \
  --match_dir data/train/match1 \
  --pred_dir pred_landing \
  --output_dir pred_landing_vis

落地/击球事件检测

tracknet/events/ 是本项目最主要的扩展模块,用于从 TrackNetV3 输出轨迹中检测关键事件。

tracknet/events/
  kinematics.py              # 速度、加速度、方向、曲率等运动学特征
  candidate_generator.py     # 规则候选生成
  detector.py                # 统一检测接口
  visual_features.py         # 视频 patch 与运动历史图特征
  bouncenet.py               # BounceNet 分类网络
  dataset.py                 # BounceNet 训练数据集

快速使用:

from tracknet.events import BounceDetector

detector = BounceDetector()
events = detector.detect_from_csv("data/test/match1/csv/1_05_02_ball.csv")

for event in events:
    print(event["frame"], event["event_type"], event["rule"])

直接传入数组时,非 512×288 视频应提供原始尺寸、FPS 和真实帧号:

detector = BounceDetector(fps=59.94, img_size=(1280, 720))
events = detector.detect(x, y, visibility, frame_ids=frame_ids)

事件类型:

事件 标识 说明
落地点 landing 球落地、停止或轨迹结束
击球点 hit 球被击打后方向或速度突变
出画面 out_of_frame 球飞出画面边缘
非事件 none BounceNet 过滤后的误检

规则检测关注高召回率,典型规则包括:

  • speed_drop:速度骤降。
  • trajectory_end:轨迹结束。
  • visibility_drop:可见性消失。
  • vy_reversal / vx_reversal:速度方向反转。
  • acceleration_peak:加速度峰值。
  • y_local_max / speed_local_max:局部极值辅助规则。

详细说明见:

半自动标注工具

标注工具用于把规则检测结果快速转成可训练标签。默认会在新文件上运行 Phase 1 检测进行候选预填充,也可以关闭自动检测进行纯手工标注。

标注工具入口为 python -m tools.labeling.launcher

单文件标注:

python -m tools.labeling.launcher --csv data/test/match1/csv/1_05_02_ball.csv

指定视频:

python -m tools.labeling.launcher \
  --csv data/test/match1/csv/1_05_02_ball.csv \
  --video data/test/match1/video/1_05_02.mp4

批量标注一个 match:

python -m tools.labeling.launcher --match_dir data/test/match1

关闭 Phase 1 自动预填充:

python -m tools.labeling.launcher \
  --csv data/test/match1/csv/1_05_02_ball.csv \
  --no-auto-detect

导出和统计:

python -m tools.labeling.launcher --export data/test/match1/labels --output training_events.csv
python -m tools.labeling.launcher --stats data/test/match1/labels

标注结果默认保存为:

data/<split>/<match>/labels/*_labels.json

详见 docs/events/labeling-tool.md

BounceNet 事件分类器

BounceNet 是二阶段事件分类器,用于过滤规则候选中的误检,并在 landing/hit/none 之间重新分类。

支持三种模式:

模式 输入 适用场景
trajectory_only 轨迹窗口 快速、无需视频,默认推荐
visual_only 视频 patch 依赖局部视觉线索
fusion 轨迹 + 视频 信息最完整,训练成本更高

训练器先读取标签 JSON 中的 csv_path/video_path 元数据,再尝试 *_labels.json*_ball.csv 的同目录约定,最后使用 --csv_dirvisual_onlyfusion 必须能解析到视频;所有模式都按 CSV 的真实 Frame 查找事件中心,并使用原视频尺寸归一化坐标。

训练:

python -m tracknet.training.bouncenet \
  --label_dir data/train/match1/labels \
  --csv_dir data/train/match1/csv \
  --mode trajectory_only \
  --epochs 100 \
  --batch_size 32 \
  --lr 1e-3 \
  --early_stopping 15 \
  --save_dir ckpts/bouncenet

视觉模式示例:

python -m tracknet.training.bouncenet \
  --label_dir data/train/match1/labels \
  --csv_dir data/train/match1/csv \
  --mode fusion \
  --device auto

推理集成:

from tracknet.events import BounceDetector

detector = BounceDetector(bouncenet_ckpt="ckpts/bouncenet/best.pt")
events = detector.detect(x, y, visibility, frames=frames, use_bouncenet=True)

训练细节见 docs/bouncenet-training.md

模型训练与评估

新包入口:

python -m tracknet.training.tracknet --help
python -m tracknet.inference.tracknet --help
python -m tracknet.evaluation.tracknet --help

训练 TrackNet:

python -m tracknet.training.tracknet \
  --model_name TrackNet \
  --seq_len 8 \
  --epochs 30 \
  --batch_size 10 \
  --bg_mode concat \
  --alpha 0.5 \
  --save_dir exp \
  --verbose

save_dir 中的 TrackNet_cur.pt 继续训练:

python -m tracknet.training.tracknet \
  --model_name TrackNet \
  --epochs 30 \
  --save_dir exp \
  --resume_training \
  --verbose

生成 InpaintNet 训练用轨迹和 mask:

python -m tools.preprocessing.generate_mask_data \
  --tracknet_file ckpts/TrackNet_best.pt \
  --batch_size 16

训练 InpaintNet:

python -m tracknet.training.tracknet \
  --model_name InpaintNet \
  --seq_len 16 \
  --epochs 300 \
  --batch_size 32 \
  --lr_scheduler StepLR \
  --mask_ratio 0.3 \
  --save_dir exp \
  --verbose

save_dir 中的 InpaintNet_cur.pt 继续训练:

python -m tracknet.training.tracknet \
  --model_name InpaintNet \
  --epochs 300 \
  --save_dir exp \
  --resume_training

评估 TrackNetV3:

python -m tools.preprocessing.generate_mask_data --tracknet_file ckpts/TrackNet_best.pt --split_list test
python -m tracknet.evaluation.tracknet --tracknet_file ckpts/TrackNet_best.pt --inpaintnet_file ckpts/InpaintNet_best.pt --save_dir eval

同时提供两个检查点时,评估会在本次运行中先执行指定 TrackNet,再把当前输出传给指定 InpaintNet;不会读取历史 predicted_csv 冒充当前 TrackNet 结果。predicted_csv 仅用于显式的离线 InpaintNet 训练数据路径。

仅评估 TrackNet:

python -m tracknet.evaluation.tracknet --tracknet_file ckpts/TrackNet_best.pt --save_dir eval

评估单个带标签视频并输出对比视频和 CSV:

python -m tracknet.evaluation.tracknet \
  --tracknet_file ckpts/TrackNet_best.pt \
  --inpaintnet_file ckpts/InpaintNet_best.pt \
  --video_file data/test/match1/video/1_05_02.mp4 \
  --save_dir eval

生成用于误差分析界面的详细预测:

python -m tracknet.evaluation.tracknet \
  --tracknet_file ckpts/TrackNet_best.pt \
  --inpaintnet_file ckpts/InpaintNet_best.pt \
  --save_dir eval \
  --output_pred

--output_pred 生成 <split>_eval_analysis_<mode>.json--output_bbox 生成按 split、drop 模式区分的 COCO 结果和 GT 元数据,避免不同评估组合复用同一全局文件。定位偏差超过 tolerance 的可见预测会同时影响 precision、recall、miss rate 和 F1,不再被当作“已召回”。

误差分析界面

tools/analysis/error_analysis.py 提供 Dash 可视化界面,用于比较不同模型或不同结果文件的逐帧误差。

通过可重复的 --eval_file 参数传入 --output_pred 生成的 JSON。一个文件可自比较,两个文件用于模型对比;程序会按真实 Frame 字段对齐。

python -m tools.analysis.error_analysis \
  --split test \
  --eval_file eval-tracknet/test_eval_analysis_weight.json \
  --eval_file eval-tracknetv3/test_eval_analysis_weight.json \
  --host 127.0.0.1

界面支持:

  • rally 级误差分布查看。
  • 逐帧预测与标签对比。
  • 不同结果文件之间的可视化比较。
  • python -m tracknet.evaluation.tracknet --output_pred 生成的 JSON 联动。
Error analysis UI

项目结构

.
├── tracknet/                        # 核心 Python 包
│   ├── models/                      # TrackNet / InpaintNet
│   ├── data/                        # 数据集与数据 I/O
│   ├── training/                    # 训练入口
│   ├── inference/                   # 推理入口
│   ├── evaluation/                  # 评估逻辑
│   ├── landing/                     # 坐标型落点预测
│   ├── events/                      # 落地/击球事件检测
│   └── visualization/               # 训练与结果可视化
├── tools/                           # 预处理、标注、Dash 分析工具
├── tests/                           # 回归测试
├── docs/                            # 使用说明与设计记录
├── assets/                          # 图片与论文
├── examples/                        # 可提交的示例标签
├── data/                            # 本地数据(不进 Git)
├── ckpts/                           # 本地模型权重(不进 Git)
├── outputs/                         # 本地预测与评估结果(不进 Git)
└── pyproject.toml                   # 项目元数据与安装配置

原始 TrackNetV3 性能参考

以下为原 TrackNetV3 在 Shuttlecock Trajectory Dataset test split 上报告的结果,用作轨迹检测主干的背景参考:

Model Accuracy Precision Recall F1 FPS
YOLOv7 57.82% 78.53% 59.96% 68.00% 34.77
TrackNetV2 94.98% 99.64% 94.56% 97.03% 27.70
TrackNetV3 97.51% 97.79% 99.33% 98.56% 25.11

本项目新增的落点预测和事件检测模块是后处理扩展,评估方式与原表中的逐帧轨迹检测指标不同,应分别查看对应输出的 metrics.jsondetail.json 或标注评估结果。

参考

License

本仓库保留原项目许可证。详见 LICENSE

About

TrackNetV3-based shuttlecock tracking and event analysis

Resources

Stars

26 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages