454 lines
14 KiB
Plaintext
454 lines
14 KiB
Plaintext
{
|
|
"cells": [
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": 1,
|
|
"metadata": {},
|
|
"outputs": [],
|
|
"source": [
|
|
"import cv2\n",
|
|
"import json\n",
|
|
"import os\n",
|
|
"import mediapipe as mp"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": 2,
|
|
"metadata": {},
|
|
"outputs": [],
|
|
"source": [
|
|
"file_path = './videos/Short_TWICE_MOONLIGHT_SUNRISE.mp4'"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"# YouTube video download (broken)"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": 3,
|
|
"metadata": {},
|
|
"outputs": [],
|
|
"source": [
|
|
"# from pytube import YouTube\n",
|
|
"\n",
|
|
"# # video_url = 'https://www.youtube.com/watch?v=E1bGxyNLIdQ'\n",
|
|
"# video_url = 'https://youtu.be/HXQyJZbs82k?si=whrz6AwRpYg5uLxl'\n",
|
|
"\n",
|
|
"# # Create a YouTube object\n",
|
|
"# yt = YouTube(video_url)\n",
|
|
"\n",
|
|
"# # Choose the stream with the desired resolution and file format (MP4)\n",
|
|
"# stream = yt.streams.filter(progressive=True, file_extension='mp4').order_by('resolution').desc().first()\n",
|
|
"\n",
|
|
"# # Define the output directory and filename\n",
|
|
"# output_dir = 'downloads'\n",
|
|
"# output_filename = yt.title + '.mp4'\n",
|
|
"\n",
|
|
"# # Download the video\n",
|
|
"# stream.download(output_path=output_dir, filename=output_filename)\n",
|
|
"\n",
|
|
"# print(f\"Downloaded: {output_dir}/{output_filename}\")"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"# Side-by-side pose estimation"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"## Utilities"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": 4,
|
|
"metadata": {},
|
|
"outputs": [],
|
|
"source": [
|
|
"\"\"\"\n",
|
|
"Concatenate frames horizontally\n",
|
|
"\"\"\"\n",
|
|
"def hconcat_resize(img_list, interpolation=cv2.INTER_CUBIC):\n",
|
|
" h_min = min(img.shape[0] for img in img_list)\n",
|
|
" im_list_resize = [cv2.resize(img, (int(img.shape[1] * h_min / img.shape[0]), h_min), interpolation=interpolation) for img in img_list]\n",
|
|
" return cv2.hconcat(im_list_resize)\n",
|
|
"\n",
|
|
"\"\"\"\n",
|
|
"JSON SERIALIZATION: Convert NormalizedLandmarkList to a nested list.\n",
|
|
"\"\"\"\n",
|
|
"def landmarks_to_list(landmarks):\n",
|
|
" if landmarks is None:\n",
|
|
" return None\n",
|
|
" landmark_list = []\n",
|
|
" for landmark in landmarks.landmark:\n",
|
|
" landmark_list.append([landmark.x, landmark.y, landmark.z])\n",
|
|
" return landmark_list\n",
|
|
"\n",
|
|
"\"\"\"\n",
|
|
"Load pose data from json file\n",
|
|
"\"\"\"\n",
|
|
"def load_pose_data(pose_data_path):\n",
|
|
" with open(pose_data_path, 'r') as json_file:\n",
|
|
" pose_data = json.load(json_file)\n",
|
|
" return pose_data\n",
|
|
"\n",
|
|
"\"\"\"\n",
|
|
"Get pose data from json file for specified frame index\n",
|
|
"\"\"\"\n",
|
|
"def get_pose_data_for_frame(pose_data, frame_number):\n",
|
|
" if frame_number < len(pose_data):\n",
|
|
" return pose_data[frame_number]\n",
|
|
" else:\n",
|
|
" return None"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"## Pre-load pose estimates for the video"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": 5,
|
|
"metadata": {},
|
|
"outputs": [],
|
|
"source": [
|
|
"\"\"\"\n",
|
|
"Pre-render pose estimation\n",
|
|
"Save baked poses into new video\n",
|
|
"\n",
|
|
"TODO: Parallelize so that loading time < video duration\n",
|
|
"\"\"\"\n",
|
|
"def precompute_and_save_pose(input_video_path, output_video_path, pose_data_path):\n",
|
|
" mp_pose = mp.solutions.pose\n",
|
|
" pose = mp_pose.Pose(min_detection_confidence=0.5, min_tracking_confidence=0.5)\n",
|
|
"\n",
|
|
" cap = cv2.VideoCapture(input_video_path)\n",
|
|
" frame_width = int(cap.get(3))\n",
|
|
" frame_height = int(cap.get(4))\n",
|
|
" out = cv2.VideoWriter(output_video_path, cv2.VideoWriter_fourcc(*'mp4v'), 30, (frame_width, frame_height))\n",
|
|
"\n",
|
|
" pose_data = []\n",
|
|
"\n",
|
|
" while cap.isOpened():\n",
|
|
" ret, frame = cap.read()\n",
|
|
" if not ret:\n",
|
|
" break\n",
|
|
"\n",
|
|
" frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)\n",
|
|
" pose_results = pose.process(frame_rgb)\n",
|
|
" mp.solutions.drawing_utils.draw_landmarks(frame, pose_results.pose_landmarks, mp.solutions.pose.POSE_CONNECTIONS)\n",
|
|
"\n",
|
|
" pose_data.append(landmarks_to_list(pose_results.pose_landmarks))\n",
|
|
" out.write(frame)\n",
|
|
"\n",
|
|
" cap.release()\n",
|
|
" out.release()\n",
|
|
"\n",
|
|
" # Save pose data to a JSON file\n",
|
|
" with open(pose_data_path, 'w') as json_file:\n",
|
|
" json.dump(pose_data, json_file)\n",
|
|
"\n",
|
|
" return pose_data\n",
|
|
"\n",
|
|
"\n",
|
|
"file_path_parts = file_path.split('.')\n",
|
|
"file_path_parts.insert(-1, 'baked_poses')\n",
|
|
"baked_file_path = '.'.join(file_path_parts)\n",
|
|
"pose_data_file_path = f\"{'.'.join(file_path_parts[:-1])}.json\"\n",
|
|
"\n",
|
|
"if os.path.exists(baked_file_path) and os.path.exists(pose_data_file_path):\n",
|
|
" file_pose_data = load_pose_data(pose_data_file_path)\n",
|
|
"else:\n",
|
|
" file_pose_data = precompute_and_save_pose(\n",
|
|
" file_path, \n",
|
|
" baked_file_path,\n",
|
|
" pose_data_file_path\n",
|
|
")"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": 24,
|
|
"metadata": {},
|
|
"outputs": [],
|
|
"source": [
|
|
"import numpy as np\n",
|
|
"from shapesimilarity import shape_similarity\n",
|
|
"# from sklearn.metrics.pairwise import cosine_similarity\n",
|
|
"# from scipy.spatial import procrustes\n",
|
|
"\n",
|
|
"weights = np.array([\n",
|
|
" # 0, # nose\n",
|
|
" # 0, # left eye (inner)\n",
|
|
" # 0, # left eye\n",
|
|
" # 0, # left eye (outer)\n",
|
|
" # 0, # right eye (inner)\n",
|
|
" # 0, # right eye\n",
|
|
" # 0, # right eye (outer)\n",
|
|
" # 0, # left ear\n",
|
|
" # 0, # right ear\n",
|
|
" # 0, # mouth (left)\n",
|
|
" # 0, # mouth (right)\n",
|
|
" # 0.2, # left shoulder\n",
|
|
" # 0.2, # right shoulder\n",
|
|
" # 0.3, # left elbow\n",
|
|
" # 0.3, # right elbow\n",
|
|
" # 0.5, # left wrist\n",
|
|
" # 0.5, # right wrist\n",
|
|
" # 0, # left pinky\n",
|
|
" # 0, # right pinky\n",
|
|
" # 0, # left index\n",
|
|
" # 0, # right index\n",
|
|
" # 0, # left thumb\n",
|
|
" # 0, # right thumb\n",
|
|
" # 0.2, # left hip\n",
|
|
" # 0.2, # right hip\n",
|
|
" # 0.3, # left knee\n",
|
|
" # 0.3, # right knee\n",
|
|
" # 0.4, # left ankle\n",
|
|
" # 0.4, # right ankle\n",
|
|
" # 0, # left heel\n",
|
|
" # 0, # right heel\n",
|
|
" # 0, # left foot index\n",
|
|
" # 0, # right foot index\n",
|
|
"\n",
|
|
" 0, # nose\n",
|
|
" 0, # left eye (inner)\n",
|
|
" 0, # left eye\n",
|
|
" 0, # left eye (outer)\n",
|
|
" 0, # right eye (inner)\n",
|
|
" 0, # right eye\n",
|
|
" 0, # right eye (outer)\n",
|
|
" 0, # left ear\n",
|
|
" 0, # right ear\n",
|
|
" 0, # mouth (left)\n",
|
|
" 0, # mouth (right)\n",
|
|
" 1, # left shoulder\n",
|
|
" 1, # right shoulder\n",
|
|
" 1, # left elbow\n",
|
|
" 1, # right elbow\n",
|
|
" 1, # left wrist\n",
|
|
" 1, # right wrist\n",
|
|
" 0, # left pinky\n",
|
|
" 0, # right pinky\n",
|
|
" 0, # left index\n",
|
|
" 0, # right index\n",
|
|
" 0, # left thumb\n",
|
|
" 0, # right thumb\n",
|
|
" 1, # left hip\n",
|
|
" 1, # right hip\n",
|
|
" 1, # left knee\n",
|
|
" 1, # right knee\n",
|
|
" 1, # left ankle\n",
|
|
" 1, # right ankle\n",
|
|
" 0, # left heel\n",
|
|
" 0, # right heel\n",
|
|
" 0, # left foot index\n",
|
|
" 0, # right foot index\n",
|
|
"]).reshape(-1, 1)\n",
|
|
"\n",
|
|
"# weights = weights / np.sum(weights)\n",
|
|
"\n",
|
|
"\n",
|
|
"def calculate_pose_similarity(pose_data1, pose_data2):\n",
|
|
" \"\"\"Calculate pose similarity using cosine similarity.\"\"\"\n",
|
|
" if pose_data1 is None or pose_data2 is None:\n",
|
|
" return None\n",
|
|
"\n",
|
|
" # Convert pose data to numpy arrays\n",
|
|
" pose1 = np.array(pose_data1)\n",
|
|
" pose2 = np.array(pose_data2)\n",
|
|
"\n",
|
|
" # Apply weights\n",
|
|
" pose1 *= weights\n",
|
|
" pose2 *= weights\n",
|
|
"\n",
|
|
" # Perform Procrustes analysis + Frechet distance\n",
|
|
" similarity = shape_similarity(pose1, pose2)\n",
|
|
" \n",
|
|
" # Calculate distance\n",
|
|
" # similarity = np.linalg.norm(pose1 - pose2)\n",
|
|
"\n",
|
|
" # Calculate cosine similarity\n",
|
|
" # similarity = cosine_similarity([pose1.flatten()], [pose2.flatten()])[0][0]\n",
|
|
" \n",
|
|
" return similarity"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"# Run pose similarity"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": 25,
|
|
"metadata": {},
|
|
"outputs": [],
|
|
"source": [
|
|
"# file_cap = cv2.VideoCapture(baked_file_path)\n",
|
|
"# webcam_cap = cv2.VideoCapture(0)\n",
|
|
"\n",
|
|
"# webcam_mp_drawing = mp.solutions.drawing_utils\n",
|
|
"# webcam_mp_pose = mp.solutions.pose\n",
|
|
"# webcam_pose = webcam_mp_pose.Pose(min_detection_confidence=0.5, min_tracking_confidence=0.5)\n",
|
|
"\n",
|
|
"# frame_number = 0\n",
|
|
"\n",
|
|
"# while file_cap.isOpened() and webcam_cap.isOpened():\n",
|
|
"# file_ret, file_frame = file_cap.read()\n",
|
|
"# webcam_ret, webcam_frame = webcam_cap.read()\n",
|
|
"\n",
|
|
"# # Break if either stream ends\n",
|
|
"# if file_frame is None or webcam_frame is None: \n",
|
|
"# break\n",
|
|
"\n",
|
|
"# webcam_frame = cv2.flip(webcam_frame, 1)\n",
|
|
"\n",
|
|
"# # Compute webcam pose estimates\n",
|
|
"# webcam_frame_rgb = cv2.cvtColor(webcam_frame, cv2.COLOR_BGR2RGB)\n",
|
|
"# webcam_pose_results = webcam_pose.process(webcam_frame_rgb)\n",
|
|
"# webcam_mp_drawing.draw_landmarks(webcam_frame, webcam_pose_results.pose_landmarks, webcam_mp_pose.POSE_CONNECTIONS)\n",
|
|
" \n",
|
|
"# # Get pose data for the current frame\n",
|
|
"# current_pose_data = get_pose_data_for_frame(file_pose_data, frame_number)\n",
|
|
" \n",
|
|
"# # Calculate pose similarity\n",
|
|
"# similarity_score = calculate_pose_similarity(current_pose_data, landmarks_to_list(webcam_pose_results.pose_landmarks)) if current_pose_data else 0\n",
|
|
"\n",
|
|
"# # Render frames\n",
|
|
"# combined_frame = hconcat_resize([file_frame, webcam_frame])\n",
|
|
" \n",
|
|
"# # Display similarity score on the frame\n",
|
|
"# if similarity_score is not None:\n",
|
|
"# similarity_text = f\"Similarity: {similarity_score:.2f}\"\n",
|
|
"# cv2.putText(combined_frame, similarity_text, (10, 30), cv2.FONT_HERSHEY_COMPLEX, 1, (0, 0, 0), 2)\n",
|
|
" \n",
|
|
"# cv2.imshow('Side-by-side', combined_frame)\n",
|
|
"\n",
|
|
"# if cv2.waitKey(1) == ord('q'):\n",
|
|
"# break\n",
|
|
"\n",
|
|
"# frame_number += 1\n",
|
|
"\n",
|
|
"# file_cap.release()\n",
|
|
"# webcam_cap.release()\n",
|
|
"# cv2.destroyAllWindows()"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"## Constant pose similarity test"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": 26,
|
|
"metadata": {},
|
|
"outputs": [],
|
|
"source": [
|
|
"import random\n",
|
|
"\n",
|
|
"file_cap = cv2.VideoCapture(baked_file_path)\n",
|
|
"webcam_cap = cv2.VideoCapture(0)\n",
|
|
"\n",
|
|
"webcam_mp_drawing = mp.solutions.drawing_utils\n",
|
|
"webcam_mp_pose = mp.solutions.pose\n",
|
|
"webcam_pose = webcam_mp_pose.Pose(min_detection_confidence=0.5, min_tracking_confidence=0.5)\n",
|
|
"\n",
|
|
"total_frames = int(file_cap.get(cv2.CAP_PROP_FRAME_COUNT))\n",
|
|
"\n",
|
|
"def get_random_frame():\n",
|
|
" random_frame_number = random.randint(0, total_frames - 1)\n",
|
|
"\n",
|
|
" file_cap.set(cv2.CAP_PROP_POS_FRAMES, random_frame_number)\n",
|
|
" file_ret, file_frame = file_cap.read()\n",
|
|
"\n",
|
|
" # Get pose data for the random frame\n",
|
|
" current_pose_data = get_pose_data_for_frame(file_pose_data, random_frame_number)\n",
|
|
" \n",
|
|
" return file_frame, current_pose_data\n",
|
|
"\n",
|
|
"file_frame, current_pose_data = get_random_frame()\n",
|
|
"\n",
|
|
"\n",
|
|
"while webcam_cap.isOpened():\n",
|
|
" webcam_ret, webcam_frame = webcam_cap.read()\n",
|
|
"\n",
|
|
" # Break if either stream ends\n",
|
|
" if webcam_frame is None: \n",
|
|
" break\n",
|
|
"\n",
|
|
" webcam_frame = cv2.flip(webcam_frame, 1)\n",
|
|
"\n",
|
|
" # Compute webcam pose estimates\n",
|
|
" webcam_frame_rgb = cv2.cvtColor(webcam_frame, cv2.COLOR_BGR2RGB)\n",
|
|
" webcam_pose_results = webcam_pose.process(webcam_frame_rgb)\n",
|
|
" webcam_mp_drawing.draw_landmarks(webcam_frame, webcam_pose_results.pose_landmarks, webcam_mp_pose.POSE_CONNECTIONS)\n",
|
|
" \n",
|
|
" # Calculate pose similarity\n",
|
|
" similarity_score = calculate_pose_similarity(current_pose_data, landmarks_to_list(webcam_pose_results.pose_landmarks)) if current_pose_data else 0\n",
|
|
"\n",
|
|
" # Render frames\n",
|
|
" combined_frame = hconcat_resize([file_frame, webcam_frame])\n",
|
|
" \n",
|
|
" # Display similarity score on the frame\n",
|
|
" if similarity_score is not None:\n",
|
|
" similarity_text = f\"Similarity: {similarity_score:.2f}\"\n",
|
|
" cv2.putText(combined_frame, similarity_text, (10, 30), cv2.FONT_HERSHEY_COMPLEX, 1, (0, 0, 0), 2)\n",
|
|
" \n",
|
|
" cv2.imshow('Side-by-side', combined_frame)\n",
|
|
"\n",
|
|
" key = cv2.waitKey(1)\n",
|
|
" if key == ord('q'):\n",
|
|
" break\n",
|
|
" elif key == ord(' '):\n",
|
|
" file_frame, current_pose_data = get_random_frame()\n",
|
|
"\n",
|
|
"file_cap.release()\n",
|
|
"webcam_cap.release()\n",
|
|
"cv2.destroyAllWindows()"
|
|
]
|
|
}
|
|
],
|
|
"metadata": {
|
|
"kernelspec": {
|
|
"display_name": "koreo-env",
|
|
"language": "python",
|
|
"name": "python3"
|
|
},
|
|
"language_info": {
|
|
"codemirror_mode": {
|
|
"name": "ipython",
|
|
"version": 3
|
|
},
|
|
"file_extension": ".py",
|
|
"mimetype": "text/x-python",
|
|
"name": "python",
|
|
"nbconvert_exporter": "python",
|
|
"pygments_lexer": "ipython3",
|
|
"version": "3.10.10"
|
|
},
|
|
"orig_nbformat": 4
|
|
},
|
|
"nbformat": 4,
|
|
"nbformat_minor": 2
|
|
}
|