Files
koreo/koreo.ipynb
2023-09-09 09:51:33 -05:00

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
}