fishxinyu's picture
download
raw
22.3 kB
"""Evaluate pose consistency of generated videos against condition pose keypoints.
Pipeline:
1. Load condition pose keypoints from <dataset-dir>/<id>_pose.npy
2. Extract pose from generated videos using DWPose
3. Normalize keypoints to [0, 1] using respective image dimensions
4. Match persons between condition and generated per frame
5. Compute AKD and PCK metrics on body keypoints (OpenPose 18-point format)
Keypoint layout (134 total, OpenPose order after DWPose wholebody):
0-17: body (nose, neck, R/L shoulder, elbow, wrist, hip, knee, ankle, eye, ear)
18-23: foot
24-91: face (68)
92-112: left hand (21)
113-133: right hand (21)
Metrics:
AKD — Average Keypoint Distance in normalized [0, 1] coord space (lower is better)
PCK@t — % of keypoints within t of GT in normalized space (higher is better)
Usage:
python eval_pose/eval.py \\
--generated-dir /path/to/generated/videos \\
--dataset-dir /path/to/condition/pose/npy \\
--dwpose-repo-path /path/to/DWPose \\
--cond-video-dir /path/to/original/condition/videos \\
--output-dir /path/to/output
"""
import argparse
import csv
import json
import os
import sys
from pathlib import Path
import av
import math
import numpy as np
import torch
import torch.nn.functional as F
from fractions import Fraction
from tqdm import tqdm
import metric
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from derange import deranged_pairing # noqa: E402
device = "cuda" if torch.cuda.is_available() else "cpu"
# First 18 keypoints = body joints (OpenPose format)
BODY_KP = list(range(18))
SCORE_THR = 0.3
PCK_THRESHOLDS = [0.05, 0.10]
eval_metrics = ["akd"] + [f"pck@{t:.2f}" for t in PCK_THRESHOLDS]
# ---------------------------------------------------------------------------
# Video I/O
# ---------------------------------------------------------------------------
def _read_video(path: Path) -> tuple[list[np.ndarray], float]:
"""Read a video into a list of [H, W, 3] uint8 numpy frames plus fps."""
with av.open(str(path)) as container:
stream = container.streams.video[0]
fps = float(stream.average_rate or stream.base_rate or 24)
frames = [f.to_ndarray(format="rgb24") for f in container.decode(video=0)]
return frames, fps
def _video_size(path: Path) -> tuple[int, int]:
"""Return (height, width) of a video without decoding all frames."""
with av.open(str(path)) as container:
s = container.streams.video[0]
return s.height, s.width
def _resize_and_center_crop(video: torch.Tensor, height: int, width: int) -> torch.Tensor:
"""Resize-to-fill then center-crop — mirrors ic_lora.py::resize_and_center_crop."""
_, _, src_h, src_w = video.shape
scale = max(height / src_h, width / src_w)
new_h = math.ceil(src_h * scale)
new_w = math.ceil(src_w * scale)
video = F.interpolate(video, size=(new_h, new_w), mode="bilinear", align_corners=False)
crop_top = (new_h - height) // 2
crop_left = (new_w - width) // 2
return video[:, :, crop_top: crop_top + height, crop_left: crop_left + width]
# ---------------------------------------------------------------------------
# DWPose model
# ---------------------------------------------------------------------------
def _load_pose_model(repo_path: Path):
"""Load DWposeDetector from a local DWPose checkout.
Mirrors compute_reference.py::load_pose_model.
"""
controlnet_root = str((repo_path / "ControlNet-v1-1-nightly").resolve())
det_ckpt = Path(controlnet_root) / "annotator" / "ckpts" / "yolox_l.onnx"
pose_ckpt = Path(controlnet_root) / "annotator" / "ckpts" / "dw-ll_ucoco_384.onnx"
for ckpt, name in [(det_ckpt, "yolox_l.onnx"), (pose_ckpt, "dw-ll_ucoco_384.onnx")]:
if not ckpt.exists():
raise FileNotFoundError(
f"Checkpoint not found: {ckpt}\n"
"Download ONNX checkpoints from https://huggingface.co/yzd-v/DWPose and "
f"place them in {ckpt.parent}/"
)
if controlnet_root not in sys.path:
sys.path.insert(0, controlnet_root)
try:
from annotator.dwpose import DWposeDetector
except ImportError as e:
raise ImportError(
f"Could not import annotator.dwpose from {controlnet_root}.\n"
f"Original error: {e}\n"
"Make sure onnxruntime is installed: pip install onnxruntime-gpu"
) from e
orig_dir = os.getcwd()
try:
os.chdir(controlnet_root)
detector = DWposeDetector()
finally:
os.chdir(orig_dir)
return detector
def _extract_pose(frames: list[np.ndarray], model) -> list[dict]:
"""Run DWPose on each frame and return raw keypoints.
Args:
frames: list of T [H, W, 3] uint8 numpy frames
model: DWposeDetector instance
Returns:
List of T dicts, each with:
'keypoints': [num_people, 134, 2] float32 — pixel coords (x, y)
'scores': [num_people, 134] float32 — confidence scores
"""
results = []
for frame in frames:
with torch.no_grad():
candidate, subset = model.pose_estimation(frame)
results.append({
"keypoints": candidate.astype(np.float32),
"scores": subset.astype(np.float32),
})
return results
# ---------------------------------------------------------------------------
# Condition NPY
# ---------------------------------------------------------------------------
def _load_condition_pose(npy_path: Path) -> list[dict]:
"""Load pre-computed condition pose keypoints from an NPY file.
Each element is a dict with:
'keypoints': [num_people, 134, 2] float32 — pixel coords (x, y)
'scores': [num_people, 134] float32
"""
data = np.load(npy_path, allow_pickle=True)
return list(data)
def _find_npy(dataset_dir: Path, stem: str) -> Path | None:
"""Find the condition NPY for a given video stem.
Tries <stem>_pose.npy first (compute_reference.py naming), then <stem>.npy.
"""
for name in (f"{stem}_pose.npy", f"{stem}.npy"):
p = dataset_dir / name
if p.exists():
return p
return None
# ---------------------------------------------------------------------------
# Person matching
# ---------------------------------------------------------------------------
def _body_centroid_normalized(kps: np.ndarray, sc: np.ndarray, h: int, w: int) -> list[tuple]:
"""Compute normalized body centroid for each detected person.
Args:
kps: [N, 134, 2] pixel coords
sc: [N, 134] scores
h, w: image height and width for normalization
Returns:
List of (centroid [2], is_valid bool) per person
"""
body_kps = kps[:, BODY_KP, :] # [N, 18, 2]
body_sc = sc[:, BODY_KP] # [N, 18]
valid = body_sc > SCORE_THR # [N, 18]
result = []
for i in range(len(kps)):
v = valid[i]
if v.any():
c = body_kps[i][v].mean(axis=0) # [2] pixel
c = c / np.array([w, h], dtype=np.float64)
result.append((c, True))
else:
result.append((np.zeros(2), False))
return result
def _match_persons(
cond_kps: np.ndarray, cond_sc: np.ndarray, cond_h: int, cond_w: int,
gen_kps: np.ndarray, gen_sc: np.ndarray, gen_h: int, gen_w: int,
) -> list[tuple[int, int]]:
"""Greedy person matching by normalized body centroid proximity.
Returns list of (cond_person_idx, gen_person_idx) pairs.
"""
cond_cents = _body_centroid_normalized(cond_kps, cond_sc, cond_h, cond_w)
gen_cents = _body_centroid_normalized(gen_kps, gen_sc, gen_h, gen_w)
matches = []
used_gen: set[int] = set()
for i, (cc, cv) in enumerate(cond_cents):
if not cv:
continue
best_j, best_dist = -1, float("inf")
for j, (gc, gv) in enumerate(gen_cents):
if not gv or j in used_gen:
continue
d = float(np.linalg.norm(cc - gc))
if d < best_dist:
best_dist, best_j = d, j
if best_j >= 0:
matches.append((i, best_j))
used_gen.add(best_j)
return matches
# ---------------------------------------------------------------------------
# Metrics
# ---------------------------------------------------------------------------
def compute_metrics(
gen_pose: list[dict],
cond_pose: list[dict],
gen_h: int,
gen_w: int,
cond_h: int,
cond_w: int,
) -> list[float]:
"""Compute AKD and PCK between generated and condition pose sequences.
Keypoints are normalized to [0, 1] by their respective image dimensions.
Only body keypoints (first 18) with score > SCORE_THR in BOTH sequences
are included. Person matching is done greedily by body centroid proximity.
Args:
gen_pose: list of T dicts with 'keypoints' [N, 134, 2] and 'scores' [N, 134]
cond_pose: same format, condition keypoints
gen_h, gen_w: generated video dimensions
cond_h, cond_w: condition video dimensions
Returns:
List of metric values in eval_metrics order.
"""
T = min(len(gen_pose), len(cond_pose))
all_pred: list[np.ndarray] = []
all_gt: list[np.ndarray] = []
for t in range(T):
gen_kps = gen_pose[t]["keypoints"] # [N_gen, 134, 2]
gen_sc = gen_pose[t]["scores"] # [N_gen, 134]
cond_kps = cond_pose[t]["keypoints"] # [N_cond, 134, 2]
cond_sc = cond_pose[t]["scores"] # [N_cond, 134]
if gen_kps.shape[0] == 0 or cond_kps.shape[0] == 0:
continue
matches = _match_persons(cond_kps, cond_sc, cond_h, cond_w,
gen_kps, gen_sc, gen_h, gen_w)
for cond_idx, gen_idx in matches:
ck = cond_kps[cond_idx, BODY_KP, :] # [18, 2]
cs = cond_sc[cond_idx, BODY_KP] # [18]
gk = gen_kps[gen_idx, BODY_KP, :] # [18, 2]
gs = gen_sc[gen_idx, BODY_KP] # [18]
valid = (cs > SCORE_THR) & (gs > SCORE_THR) # [18]
if not valid.any():
continue
# Normalize to [0, 1]: x / W, y / H
ck_norm = ck.copy().astype(np.float64)
ck_norm[:, 0] /= cond_w
ck_norm[:, 1] /= cond_h
gk_norm = gk.copy().astype(np.float64)
gk_norm[:, 0] /= gen_w
gk_norm[:, 1] /= gen_h
all_pred.append(gk_norm[valid]) # [K_valid, 2]
all_gt.append(ck_norm[valid])
if not all_pred:
return [float("nan")] * len(eval_metrics)
pred = np.concatenate(all_pred, axis=0) # [N_total, 2]
gt = np.concatenate(all_gt, axis=0)
results = [metric.akd(pred, gt)]
for thr in PCK_THRESHOLDS:
results.append(metric.pck(pred, gt, threshold=thr))
return results
# ---------------------------------------------------------------------------
# Side-by-side visualization
# ---------------------------------------------------------------------------
def _render_pose_frame(keypoints: np.ndarray, scores: np.ndarray, h: int, w: int) -> np.ndarray:
"""Render DWPose skeleton onto a black canvas [H, W, 3] uint8.
Mirrors compute_pose_reference() in compute_reference.py. annotator.dwpose must
already be on sys.path (guaranteed after _load_pose_model() runs).
"""
from annotator.dwpose import draw_pose
nums = keypoints.shape[0]
if nums == 0:
return np.zeros((h, w, 3), dtype=np.uint8)
candidate_norm = keypoints.copy().astype(float)
candidate_norm[..., 0] /= float(w)
candidate_norm[..., 1] /= float(h)
body = candidate_norm[:, :18].reshape(nums * 18, keypoints.shape[2])
score = scores[:, :18].copy()
for i in range(len(score)):
for j in range(len(score[i])):
score[i][j] = int(18 * i + j) if score[i][j] > 0.3 else -1
candidate_norm[scores < 0.3] = -1
faces = candidate_norm[:, 24:92]
hands = np.vstack([candidate_norm[:, 92:113], candidate_norm[:, 113:]])
pose = dict(bodies=dict(candidate=body, subset=score), hands=hands, faces=faces)
return draw_pose(pose, h, w)
def _save_side_by_side_video(
gen_frames: list[np.ndarray],
gen_pose: list[dict],
cond_vid_path: Path,
gen_h: int,
gen_w: int,
output_path: Path,
fps: float,
) -> None:
"""Write a 3-panel video: generated RGB | condition video (cropped) | generated pose skeleton."""
cond_video, _ = _read_video(cond_vid_path) # list of [H, W, 3] uint8
cond_tensor = torch.from_numpy(
np.stack(cond_video, axis=0)
).float().div(255.0).permute(0, 3, 1, 2) # [T, C, H, W]
cond_tensor = _resize_and_center_crop(cond_tensor, gen_h, gen_w)
cond_frames = (cond_tensor.permute(0, 2, 3, 1).numpy() * 255.0).astype(np.uint8) # [T, H, W, 3]
T = min(len(gen_frames), len(gen_pose), len(cond_frames))
panels = []
for t in range(T):
p1 = gen_frames[t]
p2 = cond_frames[t]
p3 = _render_pose_frame(gen_pose[t]["keypoints"], gen_pose[t]["scores"], gen_h, gen_w)
panels.append(np.concatenate([p1, p2, p3], axis=1))
output_path.parent.mkdir(parents=True, exist_ok=True)
with av.open(str(output_path), mode="w") as container:
stream = container.add_stream("libx264", rate=Fraction(fps).limit_denominator(1001))
stream.width = panels[0].shape[1]
stream.height = panels[0].shape[0]
stream.pix_fmt = "yuv420p"
for frame_np in panels:
av_frame = av.VideoFrame.from_ndarray(frame_np, format="rgb24")
for packet in stream.encode(av_frame):
container.mux(packet)
for packet in stream.encode():
container.mux(packet)
print(f" -> {output_path}")
# ---------------------------------------------------------------------------
# Main evaluation loop
# ---------------------------------------------------------------------------
def evaluate_pose(
generated_dir: Path,
dataset_dir: Path,
dwpose_repo_path: Path,
cond_video_dir: Path | None = None,
output_dir: Path | None = None,
show_video: bool = False,
shuffle_refs: bool = False,
) -> None:
print(f"Loading DWPose model on {device}...")
model = _load_pose_model(dwpose_repo_path)
if not dataset_dir.exists():
raise FileNotFoundError(f"Dataset (NPY) directory not found: {dataset_dir}")
cond_vid_dir = cond_video_dir or dataset_dir
generated_files = {f.stem: f for f in sorted(generated_dir.glob("*.mp4"))}
# Build NPY index: strip optional _pose suffix to get sample ID
npy_index: dict[str, Path] = {}
for f in sorted(dataset_dir.glob("*.npy")):
sid = f.stem.removesuffix("_pose")
npy_index[sid] = f
sample_ids = sorted(set(generated_files) & set(npy_index))
if not sample_ids:
print(f"No matching samples between {generated_dir} and {dataset_dir}")
return
# DWPose on the generated video is by far the dominant cost and is independent of
# which condition it is scored against, so the shuffled chance level costs only a
# second compute_metrics call on keypoints already in memory.
pairing = deranged_pairing(sample_ids) if shuffle_refs else {}
if shuffle_refs and not pairing:
print(f" Only {len(sample_ids)} sample(s); no derangement exists, skipping chance level.")
metric_names = list(eval_metrics) + ([f"{m}_chance" for m in eval_metrics] if pairing else [])
def _cond_dims(sid: str, fallback: tuple[int, int]) -> tuple[int, int]:
path = cond_vid_dir / f"{sid}.mp4"
if path.exists():
return _video_size(path)
return fallback
print(f"Evaluating {len(sample_ids)} sample(s)" + (" (+ shuffled chance)" if pairing else "") + "...")
per_sample: dict[str, dict] = {}
results_all: list[list[float]] = []
for i, sid in enumerate(tqdm(sample_ids)):
print(f" [{i + 1}/{len(sample_ids)}] {sid}", end="", flush=True)
gen_path = generated_files[sid]
npy_path = npy_index[sid]
# Load generated video and extract pose
gen_frames, gen_fps = _read_video(gen_path)
gen_h, gen_w = gen_frames[0].shape[:2]
gen_pose = _extract_pose(gen_frames, model)
# Load condition keypoints
cond_pose = _load_condition_pose(npy_path)
# Get condition video dimensions
cond_vid_path = cond_vid_dir / f"{sid}.mp4"
if cond_vid_path.exists():
cond_h, cond_w = _video_size(cond_vid_path)
else:
print(f"\n [warn] condition video not found at {cond_vid_path}; using generated dims for normalization")
cond_h, cond_w = gen_h, gen_w
sample_metrics = compute_metrics(gen_pose, cond_pose, gen_h, gen_w, cond_h, cond_w)
if pairing:
other = pairing[sid]
other_pose = _load_condition_pose(npy_index[other])
other_h, other_w = _cond_dims(other, (cond_h, cond_w))
sample_metrics = sample_metrics + compute_metrics(
gen_pose, other_pose, gen_h, gen_w, other_h, other_w
)
results_all.append(sample_metrics)
per_sample[sid] = dict(zip(metric_names, sample_metrics))
if show_video and cond_vid_path.exists():
vis_dir = output_dir / "vis" if output_dir else generated_dir / "vis"
_save_side_by_side_video(
gen_frames, gen_pose, cond_vid_path,
gen_h, gen_w,
vis_dir / f"{sid}.mp4", gen_fps,
)
metric_str = ", ".join(f"{n}={v:.4f}" for n, v in zip(metric_names, sample_metrics))
print(f" {metric_str}")
import numpy as _np
final_results = _np.array(results_all)
final_results_mean = _np.nanmean(final_results, axis=0)
print(f"\n{'=' * 50}")
print(f"Mean metrics over {len(sample_ids)} sample(s):")
for mname, mval in zip(metric_names, final_results_mean):
print(f" {mname}: {mval:.6f}")
if pairing:
means = dict(zip(metric_names, final_results_mean.tolist()))
for m in eval_metrics:
# akd is an error (lower better); pck is an accuracy (higher better).
margin = (means[f"{m}_chance"] - means[m]) if m == "akd" else (means[m] - means[f"{m}_chance"])
print(f" {m}: signal vs chance = {margin:+.6f} [{'OK' if margin > 0 else 'DEGENERATE'}]")
if output_dir is not None:
output_dir.mkdir(parents=True, exist_ok=True)
json_path = output_dir / "results.json"
result = {
"mean": {k: round(v, 4) for k, v in zip(metric_names, final_results_mean.tolist())},
"per_sample": {
sid: {m: round(v, 4) for m, v in vals.items()}
for sid, vals in per_sample.items()
},
}
with open(json_path, "w") as f:
json.dump(result, f, indent=2)
csv_path = output_dir / "results.csv"
with open(csv_path, "w", newline="") as f:
writer = csv.writer(f)
writer.writerow(["sample_id"] + metric_names)
writer.writerow(["average"] + [round(v, 4) for v in final_results_mean.tolist()])
for sid in sample_ids:
writer.writerow([sid] + [round(per_sample[sid][m], 4) for m in metric_names])
print(f"Results saved to {output_dir}")
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def main() -> None:
parser = argparse.ArgumentParser(
description="Evaluate pose consistency of generated videos against condition pose keypoints.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument(
"--generated-dir",
required=True,
help="Directory containing generated .mp4 files (one per sample ID).",
)
parser.add_argument(
"--dataset-dir",
required=True,
help=(
"Directory containing condition pose NPY files. "
"Expects files named <id>_pose.npy or <id>.npy whose stems match the generated video stems."
),
)
parser.add_argument(
"--dwpose-repo-path",
required=True,
help="Path to a local clone of DWPose (github.com/IDEA-Research/DWPose).",
)
parser.add_argument(
"--cond-video-dir",
default=None,
help=(
"Directory containing original condition videos (<id>.mp4) used to determine "
"condition keypoint resolution for normalization. "
"Defaults to --dataset-dir; if no video is found there, falls back to generated video dimensions."
),
)
parser.add_argument(
"--output-dir",
default=None,
help="Optional folder to write results.json and results.csv.",
)
parser.add_argument(
"--show-video",
action="store_true",
help=(
"Save a side-by-side comparison video for each sample: "
"generated RGB | condition pose on frame | generated pose on frame. "
"Written to <output-dir>/vis/<id>.mp4 (or <generated-dir>/vis/<id>.mp4). "
"Requires Pillow (pip install Pillow)."
),
)
parser.add_argument(
"--shuffle-refs",
action="store_true",
help=(
"Also score every video against a different sample's condition keypoints and "
"report it as <metric>_chance. Reuses the DWPose extraction, so it costs one "
"extra compute_metrics call per sample."
),
)
args = parser.parse_args()
evaluate_pose(
generated_dir=Path(args.generated_dir),
dataset_dir=Path(args.dataset_dir),
dwpose_repo_path=Path(args.dwpose_repo_path),
cond_video_dir=Path(args.cond_video_dir) if args.cond_video_dir else None,
output_dir=Path(args.output_dir) if args.output_dir else None,
show_video=args.show_video,
shuffle_refs=args.shuffle_refs,
)
if __name__ == "__main__":
main()

Xet Storage Details

Size:
22.3 kB
·
Xet hash:
860f7b3eafbe4d9ec5541ec34282ab56b248ae4c089d190c6bb769090c679bd2

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.