Download scripts/prepare_dataset.py from minhy112/FallKLTN: direct link, hf CLI and curl.
- Browser
- Download file 4.41 kB
-
https://huggingface.co/minhy112/FallKLTN/resolve/main/scripts/prepare_dataset.py
- Command line
-
hf download hf://minhy112/FallKLTN/scripts/prepare_dataset.py
-
curl -L -o prepare_dataset.py https://huggingface.co/minhy112/FallKLTN/resolve/main/scripts/prepare_dataset.py
4.41 kB
| #!/usr/bin/env python3 | |
| from __future__ import annotations | |
| import argparse | |
| from pathlib import Path | |
| import numpy as np | |
| import pandas as pd | |
| from tqdm import tqdm | |
| from fall_detection.config import load_config | |
| from fall_detection.pose_extractor import extract_video, pad_window | |
| VIDEO_EXTENSIONS = {".mp4", ".avi", ".mov", ".mkv"} | |
| def load_urfd_annotations(video_dir: Path) -> dict[str, dict[int, int]]: | |
| annotations: dict[str, dict[int, int]] = {} | |
| for path in (video_dir / "annotations").glob("*.csv"): | |
| frame = pd.read_csv(path, header=None, usecols=[0, 1, 2]) | |
| frame.columns = ["sequence", "frame", "posture"] | |
| for sequence, rows in frame.groupby("sequence"): | |
| annotations[str(sequence)] = dict( | |
| zip(rows["frame"].astype(int), rows["posture"].astype(int)) | |
| ) | |
| return annotations | |
| def window_label( | |
| class_name: str, | |
| sequence_name: str, | |
| frame_numbers: np.ndarray, | |
| annotations: dict[str, dict[int, int]], | |
| ) -> int: | |
| if class_name == "normal": | |
| return 0 | |
| if sequence_name not in annotations: | |
| return 1 | |
| mapping = annotations[sequence_name] | |
| postures = np.asarray([mapping.get(int(frame), -1) for frame in frame_numbers]) | |
| # 0 is the falling transition and 1 is lying after the fall. | |
| return int(np.mean(postures >= 0) >= 0.15) | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description="Extract MediaPipe pose windows") | |
| parser.add_argument("--config", default="configs/default.yaml") | |
| parser.add_argument("--input", type=Path) | |
| parser.add_argument("--output", type=Path) | |
| args = parser.parse_args() | |
| config = load_config(args.config) | |
| data_config = config["data"] | |
| video_dir = args.input or Path(data_config["video_dir"]) | |
| output = args.output or Path(data_config["dataset_path"]) | |
| length = int(data_config["sequence_length"]) | |
| stride = int(data_config["stride"]) | |
| min_valid_ratio = float(data_config["min_valid_pose_ratio"]) | |
| target_fps = float(data_config["target_fps"]) | |
| annotations = load_urfd_annotations(video_dir) | |
| videos = sorted( | |
| path for path in video_dir.glob("*/*/*") if path.suffix.lower() in VIDEO_EXTENSIONS | |
| ) | |
| if not videos: | |
| raise SystemExit( | |
| f"No videos found in {video_dir}/{{fall,normal}}/<group>/. " | |
| "Run scripts/download_urfd.py first." | |
| ) | |
| poses: list[np.ndarray] = [] | |
| labels: list[int] = [] | |
| groups: list[str] = [] | |
| sources: list[str] = [] | |
| skipped = 0 | |
| for video_path in tqdm(videos, desc="Extracting poses"): | |
| class_name = video_path.relative_to(video_dir).parts[0] | |
| if class_name not in {"fall", "normal"}: | |
| continue | |
| group = video_path.parent.name | |
| sequence_name = group | |
| extracted = extract_video(video_path, target_fps=target_fps) | |
| max_start = max(len(extracted.poses) - length, 0) | |
| starts = list(range(0, max_start + 1, stride)) or [0] | |
| if starts[-1] != max_start: | |
| starts.append(max_start) | |
| for start in starts: | |
| stop = start + length | |
| window = pad_window(extracted.poses[start:stop], length) | |
| frame_numbers = extracted.frame_numbers[start:stop] | |
| valid_ratio = float(np.mean(window[:, :, 3].max(axis=1) > 0)) | |
| if valid_ratio < min_valid_ratio: | |
| skipped += 1 | |
| continue | |
| poses.append(window) | |
| labels.append( | |
| window_label( | |
| class_name, sequence_name, frame_numbers, annotations | |
| ) | |
| ) | |
| groups.append(group) | |
| sources.append(f"{video_path.as_posix()}#sample={start}:{stop}") | |
| if not poses: | |
| raise SystemExit("All windows were skipped because pose detection failed") | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| np.savez_compressed( | |
| output, | |
| poses=np.stack(poses).astype(np.float32), | |
| labels=np.asarray(labels, dtype=np.int64), | |
| groups=np.asarray(groups), | |
| sources=np.asarray(sources), | |
| ) | |
| unique, counts = np.unique(labels, return_counts=True) | |
| distribution = dict(zip(unique.tolist(), counts.tolist())) | |
| print(f"Saved {len(poses)} windows to {output}") | |
| print(f"Class distribution {distribution}; skipped {skipped} low-visibility windows") | |
| if __name__ == "__main__": | |
| main() | |