Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Binary file added .DS_Store
Binary file not shown.
3 changes: 2 additions & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ figure/
__pycache__/
utils/__pycache__/
prediction/
videos/
test_prediction/
dataset_test.py
fps_exp.py
Expand All @@ -15,4 +16,4 @@ temp.py
model_compare.py
predict_old.py
*.sh
*.mp4
*.mp4
7 changes: 7 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,13 @@ scenarios to strengthen the network’s robustness. Given that a shuttlecock can
python predict.py --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
```

### Experiment logging
* `predict.py` appends a row to `experiments/experiment_log.csv` after each successful inference run.
```
python predict.py --video_file videos/test-okimoto.mp4 --tracknet_file ckpts/TrackNet_best.pt --inpaintnet_file ckpts/InpaintNet_best.pt --save_dir prediction
```
* Use `--no_experiment_log` to disable automatic logging for a run.

## Training
### 1. Prepare Dataset
* Download [Shuttlecock Trajectory Dataset](https://hackmd.io/Nf8Rh1NrSrqNUzmO0sQKZw)
Expand Down
5 changes: 3 additions & 2 deletions dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -765,8 +765,9 @@ def __gen_median__(self, max_sample_num, video_range):
else:
sample_step = 1

frame_indices = range(start_frame, end_frame, sample_step)
frame_list = []
for i in range(start_frame, end_frame, sample_step):
for i in tqdm(frame_indices, desc='median_s'):
self.cap.set(cv2.CAP_PROP_POS_FRAMES, i)
success, frame = self.cap.read()
if not success:
Expand Down Expand Up @@ -811,4 +812,4 @@ def __process__(self, imgs):
frames /= 255.
return frames



9 changes: 9 additions & 0 deletions experiments/experiment_log.csv
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
date,id,video_info,angle,background,preprocessing,resolution,fps,duration_s,bitrate,runtime_s,model_load_s,median_s,inference_s,write_csv_s,video_merge_s,result_notes,input_video,memo,accuracy,command
2026-06-27,1,BWFの試合映像 沖本,俯瞰,グリーンマット,なし,854x480,24,55,478421,,,,,,,,test-okimoto.mp4,,,
2026-06-27,3,全国選抜2024 山田,俯瞰,普通のコート,なし,1920x1080,60,40,14062749,1164,,,,,,,yamada-senbatu-crop.mp4,,,
2026-06-27,4,全国選抜2024山田,俯瞰,普通のコート,あり,1280x720,30,40,968864,,,,,,,,formated-yamada-senbatu-crop-720p30.mp4,,,
2026-06-28,5,国体2025 長束,右斜め上,グリーンマット,あり,1280x720,30,30,2825850,,,,,,,,,,
2026-06-28,6,総合2016 エンワタ,ギリ俯瞰,グリーンマット,あり,1280x720,30,93,1149310,,,,,,,,,,
2026-06-28,7,神々の遊び 3v3,横、手ブレ,選手の雑踏,あり,1280x720,30,30,2825850,,,,,,,,,,
2026-06-28,8,宮大体育館 辰樹史弥,左下,体育館の風景,あり,1280x720,30,358,1328318,,,,,,,,,,
1782803569.0459,9,,,,,1920x1080,25,46.16,2057342,844.959,0.287,272.909,559.084,0.051,12.467,,okimoto-us.mp4,,,predict.py --video_file videos/okimoto-us.mp4 --tracknet_file ckpts/TrackNet_best.pt --inpaintnet_file ckpts/InpaintNet_best.pt --save_dir prediction --batch_size 4 --output_video --large_video
67 changes: 53 additions & 14 deletions predict.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,17 @@
import os
import argparse
import sys
import time
import numpy as np
import cv2
from tqdm import tqdm

import torch
from torch.utils.data import DataLoader

from test import predict_location, get_ensemble_weight, generate_inpaint_mask
from dataset import Shuttlecock_Trajectory_Dataset, Video_IterableDataset
from utils.experiment_log import append_prediction_log
from utils.general import *


Expand Down Expand Up @@ -81,7 +85,16 @@ def predict(indices, y_pred=None, c_pred=None, img_scaler=(1, 1)):
parser.add_argument('--large_video', action='store_true', default=False, help='whether to process large video')
parser.add_argument('--output_video', action='store_true', default=False, help='whether to output video with predicted trajectory')
parser.add_argument('--traj_len', type=int, default=8, help='length of trajectory to draw on video')
parser.add_argument('--no_experiment_log', action='store_true', default=False, help='disable automatic experiment logging')
args = parser.parse_args()
run_start_time = time.perf_counter()
stage_times = {
'model_load_s': 0.0,
'median_s': 0.0,
'inference_s': 0.0,
'write_csv_s': 0.0,
'video_merge_s': 0.0,
}

num_workers = args.batch_size if args.batch_size <= 16 else 16
video_file = args.video_file
Expand All @@ -95,22 +108,27 @@ def predict(indices, y_pred=None, c_pred=None, img_scaler=(1, 1)):
os.makedirs(args.save_dir)

# Load model
tracknet_ckpt = torch.load(args.tracknet_file)
stage_start_time = time.perf_counter()
device = torch.device("mps" if torch.backends.mps.is_available() else "cpu")

tracknet_ckpt = torch.load(args.tracknet_file, map_location=device)
tracknet_seq_len = tracknet_ckpt['param_dict']['seq_len']
bg_mode = tracknet_ckpt['param_dict']['bg_mode']
tracknet = get_model('TrackNet', tracknet_seq_len, bg_mode).cuda()
tracknet = get_model('TrackNet', tracknet_seq_len, bg_mode).to(device)
tracknet.load_state_dict(tracknet_ckpt['model'])

if args.inpaintnet_file:
inpaintnet_ckpt = torch.load(args.inpaintnet_file)
inpaintnet_ckpt = torch.load(args.inpaintnet_file, map_location=device)
inpaintnet_seq_len = inpaintnet_ckpt['param_dict']['seq_len']
inpaintnet = get_model('InpaintNet').cuda()
inpaintnet = get_model('InpaintNet').to(device)
inpaintnet.load_state_dict(inpaintnet_ckpt['model'])
else:
inpaintnet = None
stage_times['model_load_s'] += time.perf_counter() - stage_start_time

cap = cv2.VideoCapture(args.video_file)
w, h = (int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)), int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)))
cap.release()
w_scaler, h_scaler = w / WIDTH, h / HEIGHT
img_scaler = (w_scaler, h_scaler)

Expand All @@ -122,8 +140,9 @@ def predict(indices, y_pred=None, c_pred=None, img_scaler=(1, 1)):
seq_len = tracknet_seq_len
if args.eval_mode == 'nonoverlap':
# Create dataset with non-overlap sampling
stage_start_time = time.perf_counter()
if large_video:
dataset = Video_IterableDataset(video_file, seq_len=seq_len, sliding_step=seq_len, bg_mode=bg_mode,
dataset = Video_IterableDataset(video_file, seq_len=seq_len, sliding_step=seq_len, bg_mode=bg_mode,
max_sample_num=args.max_sample_num, video_range=video_range)
data_loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=False, drop_last=False)
print(f'Video length: {dataset.video_len}')
Expand All @@ -133,18 +152,22 @@ def predict(indices, y_pred=None, c_pred=None, img_scaler=(1, 1)):
dataset = Shuttlecock_Trajectory_Dataset(seq_len=seq_len, sliding_step=seq_len, data_mode='heatmap', bg_mode=bg_mode,
frame_arr=np.array(frame_list)[:, :, :, ::-1], padding=True)
data_loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=False, num_workers=num_workers, drop_last=False)
stage_times['median_s'] += time.perf_counter() - stage_start_time

for step, (i, x) in enumerate(tqdm(data_loader)):
x = x.float().cuda()
stage_start_time = time.perf_counter()
for step, (i, x) in enumerate(tqdm(data_loader, desc='tracknet_inference_s')):
x = x.float().to(device)
with torch.no_grad():
y_pred = tracknet(x).detach().cpu()

# Predict
tmp_pred = predict(i, y_pred=y_pred, img_scaler=img_scaler)
for key in tmp_pred.keys():
tracknet_pred_dict[key].extend(tmp_pred[key])
stage_times['inference_s'] += time.perf_counter() - stage_start_time
else:
# Create dataset with overlap sampling for temporal ensemble
stage_start_time = time.perf_counter()
if large_video:
dataset = Video_IterableDataset(video_file, seq_len=seq_len, sliding_step=1, bg_mode=bg_mode,
max_sample_num=args.max_sample_num, video_range=video_range)
Expand All @@ -159,6 +182,7 @@ def predict(indices, y_pred=None, c_pred=None, img_scaler=(1, 1)):
frame_arr=np.array(frame_list)[:, :, :, ::-1])
data_loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=False, num_workers=num_workers, drop_last=False)
video_len = len(frame_list)
stage_times['median_s'] += time.perf_counter() - stage_start_time

# Init prediction buffer params
num_sample, sample_count = video_len-seq_len+1, 0
Expand All @@ -167,8 +191,9 @@ def predict(indices, y_pred=None, c_pred=None, img_scaler=(1, 1)):
frame_i = torch.arange(seq_len-1, -1, -1) # [7, 6, 5, 4, 3, 2, 1, 0]
y_pred_buffer = torch.zeros((buffer_size, seq_len, HEIGHT, WIDTH), dtype=torch.float32)
weight = get_ensemble_weight(seq_len, args.eval_mode)
for step, (i, x) in enumerate(tqdm(data_loader)):
x = x.float().cuda()
stage_start_time = time.perf_counter()
for step, (i, x) in enumerate(tqdm(data_loader, desc='tracknet_inference_s')):
x = x.float().to(device)
b_size, seq_len = i.shape[0], i.shape[1]
with torch.no_grad():
y_pred = tracknet(x).detach().cpu()
Expand Down Expand Up @@ -207,6 +232,7 @@ def predict(indices, y_pred=None, c_pred=None, img_scaler=(1, 1)):

# Update buffer, keep last predictions for ensemble in next iteration
y_pred_buffer = y_pred_buffer[-buffer_size:]
stage_times['inference_s'] += time.perf_counter() - stage_start_time

#assert video_len == len(tracknet_pred_dict['Frame']), 'Prediction length mismatch'
# Test on TrackNetV3 (TrackNet + InpaintNet)
Expand All @@ -221,10 +247,11 @@ def predict(indices, y_pred=None, c_pred=None, img_scaler=(1, 1)):
dataset = Shuttlecock_Trajectory_Dataset(seq_len=seq_len, sliding_step=seq_len, data_mode='coordinate', pred_dict=tracknet_pred_dict, padding=True)
data_loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=False, num_workers=num_workers, drop_last=False)

for step, (i, coor_pred, inpaint_mask) in enumerate(tqdm(data_loader)):
stage_start_time = time.perf_counter()
for step, (i, coor_pred, inpaint_mask) in enumerate(tqdm(data_loader, desc='inpaintnet_inference_s')):
coor_pred, inpaint_mask = coor_pred.float(), inpaint_mask.float()
with torch.no_grad():
coor_inpaint = inpaintnet(coor_pred.cuda(), inpaint_mask.cuda()).detach().cpu()
coor_inpaint = inpaintnet(coor_pred.to(device), inpaint_mask.to(device)).detach().cpu()
coor_inpaint = coor_inpaint * inpaint_mask + coor_pred * (1-inpaint_mask) # replace predicted coordinates with inpainted coordinates

# Thresholding
Expand All @@ -235,6 +262,7 @@ def predict(indices, y_pred=None, c_pred=None, img_scaler=(1, 1)):
tmp_pred = predict(i, c_pred=coor_inpaint, img_scaler=img_scaler)
for key in tmp_pred.keys():
inpaint_pred_dict[key].extend(tmp_pred[key])
stage_times['inference_s'] += time.perf_counter() - stage_start_time

else:
# Create dataset with overlap sampling for temporal ensemble
Expand All @@ -249,11 +277,12 @@ def predict(indices, y_pred=None, c_pred=None, img_scaler=(1, 1)):
frame_i = torch.arange(seq_len-1, -1, -1) # [7, 6, 5, 4, 3, 2, 1, 0]
coor_inpaint_buffer = torch.zeros((buffer_size, seq_len, 2), dtype=torch.float32)

for step, (i, coor_pred, inpaint_mask) in enumerate(tqdm(data_loader)):
stage_start_time = time.perf_counter()
for step, (i, coor_pred, inpaint_mask) in enumerate(tqdm(data_loader, desc='inpaintnet_inference_s')):
coor_pred, inpaint_mask = coor_pred.float(), inpaint_mask.float()
b_size = i.shape[0]
with torch.no_grad():
coor_inpaint = inpaintnet(coor_pred.cuda(), inpaint_mask.cuda()).detach().cpu()
coor_inpaint = inpaintnet(coor_pred.to(device), inpaint_mask.to(device)).detach().cpu()
coor_inpaint = coor_inpaint * inpaint_mask + coor_pred * (1-inpaint_mask)

# Thresholding
Expand Down Expand Up @@ -299,14 +328,24 @@ def predict(indices, y_pred=None, c_pred=None, img_scaler=(1, 1)):

# Update buffer, keep last predictions for ensemble in next iteration
coor_inpaint_buffer = coor_inpaint_buffer[-buffer_size:]
stage_times['inference_s'] += time.perf_counter() - stage_start_time


# Write csv file
pred_dict = inpaint_pred_dict if inpaintnet is not None else tracknet_pred_dict
stage_start_time = time.perf_counter()
write_pred_csv(pred_dict, save_file=out_csv_file)
stage_times['write_csv_s'] += time.perf_counter() - stage_start_time

# Write video with predicted coordinates
if args.output_video:
stage_start_time = time.perf_counter()
write_pred_video(video_file, pred_dict, save_file=out_video_file, traj_len=args.traj_len)
stage_times['video_merge_s'] += time.perf_counter() - stage_start_time

runtime_s = time.perf_counter() - run_start_time

if not args.no_experiment_log:
append_prediction_log(args, runtime_s, stage_times, sys.argv)

print('Done.')
print('Done.')
8 changes: 5 additions & 3 deletions requirements.txt
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
dash==2.5.1
numpy==1.22.4
opencv_python==4.4.0.46
opencv_python==4.13.0.92
pandas==2.0.0
Pillow==10.0.0
plotly==5.8.2
torch==1.10.0
parse
torch==2.4.1
parse
tqdm==4.68.3
pycocotools==2.0.7
122 changes: 122 additions & 0 deletions utils/experiment_log.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,122 @@
import csv
import os
import shlex
import time
from pathlib import Path

import cv2


EXPERIMENT_LOG_FILE = Path("experiments/experiment_log.csv")
EXPERIMENT_LOG_FIELDS = [
"date",
"id",
"video_info",
"angle",
"background",
"preprocessing",
"resolution",
"fps",
"duration_s",
"bitrate",
"runtime_s",
"model_load_s",
"median_s",
"inference_s",
"write_csv_s",
"video_merge_s",
"result_notes",
"input_video",
"memo",
"accuracy",
"command",
]


def format_seconds(seconds):
return f"{seconds:.3f}"


def read_experiment_rows(log_file):
if not log_file.exists():
return []

with log_file.open("r", encoding="utf-8", newline="") as f:
return list(csv.DictReader(f))


def next_experiment_id(rows):
ids = []
for row in rows:
try:
ids.append(int(row.get("id", "")))
except ValueError:
continue
return max(ids, default=0) + 1


def get_video_metadata(video_file):
cap = cv2.VideoCapture(video_file)
if not cap.isOpened():
return {
"resolution": "",
"fps": "",
"duration_s": "",
"bitrate": "",
}

width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
fps = cap.get(cv2.CAP_PROP_FPS)
frame_count = cap.get(cv2.CAP_PROP_FRAME_COUNT)
cap.release()

duration_s = frame_count / fps if fps and fps > 0 else 0
file_size_bits = os.path.getsize(video_file) * 8 if os.path.exists(video_file) else 0
bitrate = int(file_size_bits / duration_s) if duration_s > 0 else ""

return {
"resolution": f"{width}x{height}" if width and height else "",
"fps": f"{fps:.3f}".rstrip("0").rstrip(".") if fps else "",
"duration_s": f"{duration_s:.3f}".rstrip("0").rstrip(".") if duration_s else "",
"bitrate": bitrate,
}


def append_prediction_log(args, runtime_s, stage_times, argv):
log_file = EXPERIMENT_LOG_FILE
metadata = get_video_metadata(args.video_file)
rows = read_experiment_rows(log_file)
row = {
"date": time.time(),
"id": next_experiment_id(rows),
"video_info": "",
"angle": "",
"background": "",
"preprocessing": "",
"resolution": metadata["resolution"],
"fps": metadata["fps"],
"duration_s": metadata["duration_s"],
"bitrate": metadata["bitrate"],
"runtime_s": format_seconds(runtime_s),
"model_load_s": format_seconds(stage_times.get('model_load_s', 0)),
"median_s": format_seconds(stage_times.get('median_s', 0)),
"inference_s": format_seconds(stage_times.get('inference_s', 0)),
"write_csv_s": format_seconds(stage_times.get('write_csv_s', 0)),
"video_merge_s": format_seconds(stage_times.get('video_merge_s', 0)),
"result_notes": "",
"input_video": os.path.basename(args.video_file),
"memo": "",
"accuracy": "",
"command": " ".join(shlex.quote(arg) for arg in argv),
}

log_file.parent.mkdir(parents=True, exist_ok=True)
exists = log_file.exists()
with log_file.open("a", encoding="utf-8", newline="") as f:
writer = csv.DictWriter(f, fieldnames=EXPERIMENT_LOG_FIELDS)
if not exists:
writer.writeheader()
writer.writerow(row)

print(f"Experiment log: appended #{row['id']} to {log_file}")
Loading