使用TrackNet跟踪网球时出现定位偏移问题求助
TrackNet网球跟踪标记偏移问题排查
我使用TrackNet模型对网球比赛视频中的球进行跟踪,发现跟踪轨迹能完美跟随球的运动,但标记圆圈并未精准落在球上,始终存在偏移。

以下是我使用的TrackNet代码:
from model import BallTrackerNet import torch import cv2 from general import postprocess from tqdm import tqdm import numpy as np import argparse from itertools import groupby from scipy.spatial import distance def read_video(path_video): """ Read video file :params path_video: path to video file :return frames: list of video frames fps: frames per second """ cap = cv2.VideoCapture(path_video) fps = int(cap.get(cv2.CAP_PROP_FPS)) frames = [] while cap.isOpened(): ret, frame = cap.read() if ret: frames.append(frame) else: break cap.release() return frames, fps def infer_model(frames, model): """ Run pretrained model on a consecutive list of frames :params frames: list of consecutive video frames model: pretrained model :return ball_track: list of detected ball points dists: list of euclidean distances between two neighbouring ball points """ height = 360 width = 640 dists = [-1]*2 ball_track = [(None,None)]*2 for num in tqdm(range(2, len(frames))): img = cv2.resize(frames[num], (width, height)) img_prev = cv2.resize(frames[num-1], (width, height)) img_preprev = cv2.resize(frames[num-2], (width, height)) imgs = np.concatenate((img, img_prev, img_preprev), axis=2) imgs = imgs.astype(np.float32)/255.0 imgs = np.rollaxis(imgs, 2, 0) inp = np.expand_dims(imgs, axis=0) out = model(torch.from_numpy(inp).float().to(device)) output = out.argmax(dim=1).detach().cpu().numpy() x_pred, y_pred = postprocess(output) ball_track.append((x_pred, y_pred)) if ball_track[-1][0] and ball_track[-2][0]: dist = distance.euclidean(ball_track[-1], ball_track[-2]) else: dist = -1 dists.append(dist) return ball_track, dists def remove_outliers(ball_track, dists, max_dist = 100): """ Remove outliers from model prediction :params ball_track: list of detected ball points dists: list of euclidean distances between two neighbouring ball points max_dist: maximum distance between two neighbouring ball points :return ball_track: list of ball points """ outliers = list(np.where(np.array(dists) > max_dist)[0]) for i in outliers: if (dists[i+1] > max_dist) | (dists[i+1] == -1): ball_track[i] = (None, None) outliers.remove(i) elif dists[i-1] == -1: ball_track[i-1] = (None, None) return ball_track def split_track(ball_track, max_gap=4, max_dist_gap=80, min_track=5): """ Split ball track into several subtracks in each of which we will perform ball interpolation. :params ball_track: list of detected ball points max_gap: maximun number of coherent None values for interpolation max_dist_gap: maximum distance at which neighboring points remain in one subtrack min_track: minimum number of frames in each subtrack :return result: list of subtrack indexes """ list_det = [0 if x[0] else 1 for x in ball_track] groups = [(k, sum(1 for _ in g)) for k, g in groupby(list_det)] cursor = 0 min_value = 0 result = [] for i, (k, l) in enumerate(groups): if (k == 1) & (i > 0) & (i < len(groups) - 1): dist = distance.euclidean(ball_track[cursor-1], ball_track[cursor+l]) if (l >=max_gap) | (dist/l > max_dist_gap): if cursor - min_value > min_track: result.append([min_value, cursor]) min_value = cursor + l - 1 cursor += l if len(list_det) - min_value > min_track: result.append([min_value, len(list_det)]) return result def interpolation(coords): """ Run ball interpolation in one subtrack :params coords: list of ball coordinates of one subtrack :return track: list of interpolated ball coordinates of one subtrack """ def nan_helper(y): return np.isnan(y), lambda z: z.nonzero()[0] x = np.array([x[0] if x[0] is not None else np.nan for x in coords]) y = np.array([x[1] if x[1] is not None else np.nan for x in coords]) nons, yy = nan_helper(x) x[nons]= np.interp(yy(nons), yy(~nons), x[~nons]) nans, xx = nan_helper(y) y[nans]= np.interp(xx(nans), xx(~nans), y[~nans]) track = [*zip(x,y)] return track def write_track(frames, ball_track, path_output_video, fps, trace=7): """ Write .avi file with detected ball tracks :params frames: list of original video frames ball_track: list of ball coordinates path_output_video: path to output video fps: frames per second trace: number of frames with detected trace """ height, width = frames[0].shape[:2] # out = cv2.VideoWriter(path_output_video, cv2.VideoWriter_fourcc(*'DIVX'), # fps, (width, height)) out = cv2.VideoWriter(path_output_video, cv2.VideoWriter_fourcc(*'mp4v'), fps, (width, height)) for num in range(len(frames)): frame = frames[num] for i in range(trace): if (num-i > 0): if ball_track[num-i][0]: x = int(ball_track[num-i][0]) y = int(ball_track[num-i][1]) frame = cv2.circle(frame, (x,y), radius=0, color=(0, 0, 255), thickness=10-i) else: break out.write(frame) out.release() if __name__ == '__main__': parser = argparse.ArgumentParser() parser.add_argument('--batch_size', type=int, default=2, help='batch size') parser.add_argument('--model_path', type=str, help='path to model') parser.add_argument('--video_path', type=str, help='path to input video') parser.add_argument('--video_out_path', type=str, help='path to output video') parser.add_argument('--extrapolation', action='store_true', help='whether to use ball track extrapolation') args = parser.parse_args() model = BallTrackerNet() device = 'cpu' model.load_state_dict(torch.load(args.model_path, map_location=device)) model = model.to(device) model.eval() frames, fps = read_video(args.video_path) ball_track, dists = infer_model(frames, model) ball_track = remove_outliers(ball_track, dists) if args.extrapolation: subtracks = split_track(ball_track) for r in subtracks: ball_subtrack = ball_track[r[0]:r[1]] ball_subtrack = interpolation(ball_subtrack) ball_track[r[0]:r[1]] = ball_subtrack write_track(frames, ball_track, args.video_out_path, fps)
可能的原因及解决方法
- 坐标未做缩放还原:代码中
infer_model将视频帧缩放到了640×360的尺寸进行推理,但postprocess返回的坐标是基于这个缩放后尺寸的,后续write_track直接用该坐标在原始尺寸的帧上绘制标记,必然会出现偏移。解决方法:在infer_model中记录原始帧的宽高,将预测的(x_pred, y_pred)按比例还原:# 在infer_model函数开头,获取原始帧尺寸 orig_h, orig_w = frames[0].shape[:2] # 得到x_pred, y_pred后,还原坐标 x_pred = x_pred * (orig_w / width) if x_pred is not None else None y_pred = y_pred * (orig_h / height) if y_pred is not None else None - postprocess函数坐标映射错误:检查
general.py中的postprocess函数,确认它是否正确将模型输出的热力图坐标转换为缩放后图像的坐标。比如模型输出的特征图尺寸如果小于640×360,是否做了正确的上采样映射。 - 通道顺序不一致:OpenCV读取的视频帧是BGR通道,而TrackNet训练时可能使用的是RGB通道。通道顺序不匹配会导致模型特征提取偏差,进而引发坐标偏移。解决方法:在缩放后添加通道转换:
img = cv2.cvtColor(cv2.resize(frames[num], (width, height)), cv2.COLOR_BGR2RGB) img_prev = cv2.cvtColor(cv2.resize(frames[num-1], (width, height)), cv2.COLOR_BGR2RGB) img_preprev = cv2.cvtColor(cv2.resize(frames[num-2], (width, height)), cv2.COLOR_BGR2RGB) - 训练数据与推理数据预处理不一致:如果使用的是预训练模型,训练时的图像预处理(如裁剪、平移、亮度调整等)与推理时的操作不同,也会导致预测偏移。可以对比训练代码的预处理步骤,确保推理时的操作完全一致。
内容的提问来源于stack exchange,提问作者Plewis
相关产品推荐
相关产品推荐

