final demo #1
Browse files- app.py +6 -15
- backend.py +184 -0
- configs/README.md +70 -0
- configs/predict_config_demo.yaml +58 -0
- predict_pose.py +529 -0
- utils/__pycache__/translation.cpython-311.pyc +0 -0
- utils/translation.py +32 -0
app.py
CHANGED
|
@@ -1,23 +1,14 @@
|
|
| 1 |
import gradio as gr
|
| 2 |
import os
|
| 3 |
os.environ["KMP_DUPLICATE_LIB_OK"]="TRUE"
|
| 4 |
-
|
| 5 |
-
from huggingface_hub import hf_hub_download
|
| 6 |
-
# ==========================================
|
| 7 |
-
# 1. DOWNLOAD THE LARGE FILE ON STARTUP
|
| 8 |
-
# ==========================================
|
| 9 |
-
print("Checking for large model file...")
|
| 10 |
-
|
| 11 |
-
# UPDATE THESE TWO LINES WITH YOUR EXACT DETAILS
|
| 12 |
-
model_path = hf_hub_download(
|
| 13 |
-
repo_id="plice13/sign-language-weights",
|
| 14 |
-
filename="best_checkpoint.pth"
|
| 15 |
-
)
|
| 16 |
-
print(f"File successfully loaded at: {model_path}")
|
| 17 |
-
# ==========================================
|
| 18 |
|
| 19 |
def process_video(input_video_path):
|
| 20 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 21 |
|
| 22 |
# Custom CSS for the dark green background and centered layout
|
| 23 |
custom_css = """
|
|
|
|
| 1 |
import gradio as gr
|
| 2 |
import os
|
| 3 |
os.environ["KMP_DUPLICATE_LIB_OK"]="TRUE"
|
| 4 |
+
from backend import process_input
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5 |
|
| 6 |
def process_video(input_video_path):
|
| 7 |
+
# Generate a translation in the backend.
|
| 8 |
+
translation = process_input(input_video_path)
|
| 9 |
+
|
| 10 |
+
# Return just the translation if no visualization is found
|
| 11 |
+
return translation, None
|
| 12 |
|
| 13 |
# Custom CSS for the dark green background and centered layout
|
| 14 |
custom_css = """
|
backend.py
ADDED
|
@@ -0,0 +1,184 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import time
|
| 3 |
+
import torch
|
| 4 |
+
import numpy as np
|
| 5 |
+
import matplotlib.pyplot as plt
|
| 6 |
+
from dotenv import load_dotenv
|
| 7 |
+
|
| 8 |
+
# Import tvého nového modelu a pomocných funkcí
|
| 9 |
+
from Uni_Sign.models import Uni_Sign
|
| 10 |
+
from Uni_Sign.datasets import load_part_kp_YTASL, YTASL_GROUP_SIZES, _fill_missing_landmarks, select_frame_indices
|
| 11 |
+
from predict_pose import create_mediapipe_models, predict_pose, load_video_cv
|
| 12 |
+
|
| 13 |
+
os.environ['KMP_DUPLICATE_LIB_OK'] = 'TRUE'
|
| 14 |
+
load_dotenv()
|
| 15 |
+
|
| 16 |
+
# Globální proměnné pro modely, ať se nenačítají při každém videu znovu
|
| 17 |
+
model = None
|
| 18 |
+
pose_models = None
|
| 19 |
+
args = None
|
| 20 |
+
|
| 21 |
+
class InferenceConfig:
|
| 22 |
+
"""Čistá konfigurace pro Uni_Sign inferenci"""
|
| 23 |
+
def __init__(self):
|
| 24 |
+
# Nezapomeň upravit cestu ke svým natrénovaným váhám
|
| 25 |
+
self.finetune = r"./Uni_Sign/unisign_model/best_checkpoint.pth"
|
| 26 |
+
self.dataset = "YTASL"
|
| 27 |
+
self.task = "SLT"
|
| 28 |
+
self.max_length = 256
|
| 29 |
+
self.normalization = "none"
|
| 30 |
+
self.layout = "pruned"
|
| 31 |
+
self.n_registers = 0
|
| 32 |
+
self.hidden_dim = 256
|
| 33 |
+
self.rgb_support = False
|
| 34 |
+
self.no_adaptive_gcn = False
|
| 35 |
+
self.register_position = "before_all"
|
| 36 |
+
self.label_smoothing = 0.2
|
| 37 |
+
self.batch_size = 1
|
| 38 |
+
|
| 39 |
+
def initialize_model():
|
| 40 |
+
"""Inicializuje Uni_Sign a Mediapipe/YOLO modely pro extrakci pose."""
|
| 41 |
+
global model, pose_models, args
|
| 42 |
+
|
| 43 |
+
if model is not None and pose_models is not None:
|
| 44 |
+
return
|
| 45 |
+
|
| 46 |
+
print("Inicializuji Uni_Sign model...")
|
| 47 |
+
args = InferenceConfig()
|
| 48 |
+
model = Uni_Sign(args=args)
|
| 49 |
+
|
| 50 |
+
# Načtení vah
|
| 51 |
+
if args.finetune and os.path.exists(args.finetune):
|
| 52 |
+
print(f"Načítám váhy z: {args.finetune}")
|
| 53 |
+
state_dict = torch.load(args.finetune, map_location='cpu')['model']
|
| 54 |
+
model.load_state_dict(state_dict, strict=False)
|
| 55 |
+
else:
|
| 56 |
+
print(f"VAROVÁNÍ: Checkpoint '{args.finetune}' nebyl nalezen!")
|
| 57 |
+
|
| 58 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 59 |
+
model.to(device)
|
| 60 |
+
model.eval()
|
| 61 |
+
|
| 62 |
+
print("Inicializuji Mediapipe a YOLO modely...")
|
| 63 |
+
pose_checkpoint_folder = 'checkpoints/pose/'
|
| 64 |
+
pose_models = create_mediapipe_models(pose_checkpoint_folder)
|
| 65 |
+
|
| 66 |
+
def process_pose_data_in_memory(pose_results, args):
|
| 67 |
+
"""
|
| 68 |
+
Nahrazuje starý 'process_single_json'.
|
| 69 |
+
Zpracovává 'cropped_keypoints' přímo ze slovníku (v paměti).
|
| 70 |
+
"""
|
| 71 |
+
raw_pose = pose_results.get('cropped_keypoints', [])
|
| 72 |
+
if not raw_pose:
|
| 73 |
+
raise ValueError("Video neobsahuje žádné detekované pose (cropped_keypoints).")
|
| 74 |
+
|
| 75 |
+
# NOVÝ KÓD: Převod numpy polí na čisté [X, Y] seznamy a ošetření chybějících částí těla
|
| 76 |
+
pose = []
|
| 77 |
+
for frame_data in raw_pose:
|
| 78 |
+
formatted_frame = {}
|
| 79 |
+
for part, expected_size in YTASL_GROUP_SIZES.items():
|
| 80 |
+
kps = frame_data.get(part, [])
|
| 81 |
+
|
| 82 |
+
# Pokud část chybí (prázdný list/array), pošleme prázdno, ať si s tím poradí _fill_missing_landmarks
|
| 83 |
+
# Pokud část chybí (prázdný list/array), pošleme prázdno
|
| 84 |
+
# Pokud část chybí (prázdný list/array), pošleme prázdno
|
| 85 |
+
# Bezpečná kontrola Numpy pole:
|
| 86 |
+
if kps is None or len(kps) == 0:
|
| 87 |
+
formatted_frame[part] = []
|
| 88 |
+
else:
|
| 89 |
+
# Neprůstřelné řešení: převedeme data (ať už jsou cokoliv) na Numpy pole,
|
| 90 |
+
# ořízneme první dva sloupce a převedeme zpět na čistý list.
|
| 91 |
+
formatted_frame[part] = np.array(kps)[:, :2].tolist()
|
| 92 |
+
pose.append(formatted_frame)
|
| 93 |
+
|
| 94 |
+
# ... Zbytek funkce pokračuje beze změny:
|
| 95 |
+
|
| 96 |
+
duration = len(pose)
|
| 97 |
+
tmp = select_frame_indices(duration, args.max_length, phase='test')
|
| 98 |
+
skeletons = [pose[i] for i in tmp]
|
| 99 |
+
|
| 100 |
+
confs = []
|
| 101 |
+
for i, skeleton in enumerate(skeletons):
|
| 102 |
+
conf = {}
|
| 103 |
+
for group_name, expected_size in YTASL_GROUP_SIZES.items():
|
| 104 |
+
_fill_missing_landmarks(
|
| 105 |
+
skeleton=skeleton,
|
| 106 |
+
conf=conf,
|
| 107 |
+
group_name=group_name,
|
| 108 |
+
expected_size=expected_size,
|
| 109 |
+
clip_name="gradio_video",
|
| 110 |
+
frame_idx=i,
|
| 111 |
+
error_group_label=f"group '{group_name}'",
|
| 112 |
+
include_size_details=False,
|
| 113 |
+
strict_key_access=True,
|
| 114 |
+
)
|
| 115 |
+
confs.append(conf)
|
| 116 |
+
|
| 117 |
+
kps_with_scores = load_part_kp_YTASL(skeletons, confs, args.normalization, args.layout)
|
| 118 |
+
|
| 119 |
+
src_input = {}
|
| 120 |
+
for key, val in kps_with_scores.items():
|
| 121 |
+
src_input[key] = val.unsqueeze(0)
|
| 122 |
+
|
| 123 |
+
seq_len = src_input['body'].shape[1]
|
| 124 |
+
mask_gen = torch.ones([seq_len]) + 7
|
| 125 |
+
src_input['attention_mask'] = (mask_gen != 0).long().unsqueeze(0)
|
| 126 |
+
src_input['src_length_batch'] = torch.LongTensor([seq_len])
|
| 127 |
+
src_input['name_batch'] = ["gradio_video"]
|
| 128 |
+
|
| 129 |
+
return src_input
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
def process_input(input_video_path):
|
| 133 |
+
"""
|
| 134 |
+
Hlavní vstupní bod pro Gradio. Spustí extrakci dat, vizualizaci
|
| 135 |
+
a inference přes Uni_Sign.
|
| 136 |
+
"""
|
| 137 |
+
try:
|
| 138 |
+
initialize_model()
|
| 139 |
+
|
| 140 |
+
print("1. Extrahuji keypoints z videa...")
|
| 141 |
+
video_frames, _ = load_video_cv(input_video_path)
|
| 142 |
+
pose_results = predict_pose(video_frames, pose_models)
|
| 143 |
+
|
| 144 |
+
print("2. Prostor pro vizualizaci...")
|
| 145 |
+
|
| 146 |
+
print("3. Přenáším data do Tenzorů pro Uni_Sign...")
|
| 147 |
+
src_input = process_pose_data_in_memory(pose_results, args)
|
| 148 |
+
|
| 149 |
+
# Přesunutí Tenzorů na stejné zařízení (GPU/CPU) jako model
|
| 150 |
+
device = next(model.parameters()).device
|
| 151 |
+
for key in ['body', 'left', 'right', 'face_all', 'attention_mask']:
|
| 152 |
+
if key in src_input:
|
| 153 |
+
# Ošetření typu dat na float
|
| 154 |
+
if key != 'attention_mask':
|
| 155 |
+
src_input[key] = src_input[key].to(device, dtype=torch.float32)
|
| 156 |
+
else:
|
| 157 |
+
src_input[key] = src_input[key].to(device)
|
| 158 |
+
|
| 159 |
+
tgt_input = {'gt_sentence': [""], 'gt_gloss': [""]}
|
| 160 |
+
|
| 161 |
+
print("4. Generuji překlad...")
|
| 162 |
+
with torch.no_grad():
|
| 163 |
+
stack_out = model(src_input, tgt_input)
|
| 164 |
+
output = model.generate(
|
| 165 |
+
stack_out,
|
| 166 |
+
max_new_tokens=100,
|
| 167 |
+
num_beams=4,
|
| 168 |
+
)
|
| 169 |
+
|
| 170 |
+
tokenizer = model.mt5_tokenizer
|
| 171 |
+
tgt_pres = tokenizer.batch_decode(output, skip_special_tokens=True)
|
| 172 |
+
|
| 173 |
+
result = tgt_pres[0].strip()
|
| 174 |
+
if not result:
|
| 175 |
+
return "Model nedokázal generovat text. Zkus jiné video."
|
| 176 |
+
|
| 177 |
+
print("5. VŠE HOTOVO!")
|
| 178 |
+
return result
|
| 179 |
+
|
| 180 |
+
except Exception as e:
|
| 181 |
+
print(f"Chyba při zpracování: {e}")
|
| 182 |
+
import traceback
|
| 183 |
+
traceback.print_exc()
|
| 184 |
+
return f"Chyba při zpracování videa: {str(e)}"
|
configs/README.md
ADDED
|
@@ -0,0 +1,70 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
## Training config
|
| 2 |
+
|
| 3 |
+
```
|
| 4 |
+
ModelArguments:
|
| 5 |
+
base_model_name: Name of the model to be downloaded from HF
|
| 6 |
+
sign_input_dim (depricated, it is calculated automatically now): Dimmension of sign features on the model input (to be changed with multimodal input)
|
| 7 |
+
hidden_dropout_prob: Dropout probability of the input layer
|
| 8 |
+
max_length: Max sequence length to be generated by the model
|
| 9 |
+
|
| 10 |
+
# TrainingArguments are overwritten by not-None arguments
|
| 11 |
+
TrainingArguments:
|
| 12 |
+
project_name: Name of wandb project
|
| 13 |
+
model_name: Name of the run name (will affect wandb and checkpoint path)
|
| 14 |
+
output_dir: Output path for checkpoints
|
| 15 |
+
resume_from_checkpoint: Path to checkpoint to be loaded. The path needs to contain model.safetensors file
|
| 16 |
+
load_only_weights: True for loading only weights to use different scheduler or optimizer. False for loading the whole checkpoint
|
| 17 |
+
# Logging and saving
|
| 18 |
+
report_to: "wandb" for wandb report. Otherwise wandb not used
|
| 19 |
+
# Debugging
|
| 20 |
+
max_train_samples: Maximum size of training dataset for debugging. "none" for default size
|
| 21 |
+
max_val_samples: Maximum size of validation dataset for debugging. "none" for default size
|
| 22 |
+
# Data processing
|
| 23 |
+
max_sequence_length: Max input sequence length (decoder context window) - data sequence is cropped by this length
|
| 24 |
+
max_token_length: Dataloader tokenizer max_length
|
| 25 |
+
skip_frames: Use only each n-th frame. "True" for each 2nd frame
|
| 26 |
+
|
| 27 |
+
SignDataArguments:
|
| 28 |
+
data_dir: Prefix of the path to data including annotation file and metadatafile. This path is joined by paths bellow.
|
| 29 |
+
annotation_path: Path to train/dev annotation file
|
| 30 |
+
visual_features: Path to train/dev metadata file
|
| 31 |
+
pose:
|
| 32 |
+
enable_input: Whether to use this modality. Our code supports only "pose" so far
|
| 33 |
+
train: Path to train metafile
|
| 34 |
+
dev: Path to validation metafile
|
| 35 |
+
|
| 36 |
+
SignModelArguments: Dimmensions per each sign-features modality
|
| 37 |
+
```
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
## Predict config
|
| 41 |
+
|
| 42 |
+
```
|
| 43 |
+
ModelArguments:
|
| 44 |
+
base_model_name: Name of the model to be downloaded from HF
|
| 45 |
+
sign_input_dim: Dimmension of sign features on the model input (to be changed with multimodal input)
|
| 46 |
+
max_length: Max sequence length to be generated by the model
|
| 47 |
+
|
| 48 |
+
EvaluationArguments:
|
| 49 |
+
output_dir: Output path for predictions
|
| 50 |
+
model_name: Only for logging purposes
|
| 51 |
+
skip_frames: Use only each n-th frame. "True" for each 2nd frame
|
| 52 |
+
# Data processing
|
| 53 |
+
split: test
|
| 54 |
+
max_sequence_length: Max input sequence length (decoder context window) - data sequence is cropped by this length
|
| 55 |
+
max_token_length: Dataloader tokenizer max_length
|
| 56 |
+
# Generation parameters
|
| 57 |
+
model_dir: Path to checkpoint to be evaluated. The path needs to contain model.safetensors file
|
| 58 |
+
# Debugging
|
| 59 |
+
max_val_samples: Maximum size of validation dataset for debugging. "none" for default size
|
| 60 |
+
|
| 61 |
+
SignDataArguments:
|
| 62 |
+
data_dir: Prefix of the path to data including annotation file and metadatafile. This path is joined by paths bellow.
|
| 63 |
+
annotation_path: Path to test annotation file
|
| 64 |
+
visual_features: Path to test metadata file
|
| 65 |
+
pose:
|
| 66 |
+
enable_input: Whether to use this modality. Our code supports only "pose" so far
|
| 67 |
+
test: Path to test metafile
|
| 68 |
+
|
| 69 |
+
SignModelArguments: Dimmensions per each sign-features modality
|
| 70 |
+
```
|
configs/predict_config_demo.yaml
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
ModelArguments:
|
| 2 |
+
base_model_name: google/t5-v1_1-base
|
| 3 |
+
hidden_dropout_prob: 0.1
|
| 4 |
+
num_beams: 5
|
| 5 |
+
max_length: 128
|
| 6 |
+
top_k: 50
|
| 7 |
+
top_p: 0.9
|
| 8 |
+
temperature: 1.0
|
| 9 |
+
length_penalty: 1.0
|
| 10 |
+
repetition_penalty: 1.0
|
| 11 |
+
early_stopping: False
|
| 12 |
+
no_repeat_ngram_size: 3
|
| 13 |
+
do_sample: False
|
| 14 |
+
|
| 15 |
+
EvaluationArguments:
|
| 16 |
+
output_dir: ./results
|
| 17 |
+
model_name: T5_YTASL_pretrain
|
| 18 |
+
skip_frames: False
|
| 19 |
+
# Data processing
|
| 20 |
+
split: test
|
| 21 |
+
max_sequence_length: 250
|
| 22 |
+
max_token_length: 128
|
| 23 |
+
# Generation parameters
|
| 24 |
+
model_dir: ./checkpoints/t5-v1_1-base/model.safetensors
|
| 25 |
+
batch_size: 1
|
| 26 |
+
# Debugging
|
| 27 |
+
max_val_samples: none
|
| 28 |
+
|
| 29 |
+
SignDataArguments:
|
| 30 |
+
data_dir: ./data
|
| 31 |
+
annotation_path:
|
| 32 |
+
train: YT.annotations.train.json
|
| 33 |
+
dev: YT.annotations.dev.json
|
| 34 |
+
test: YT.annotations.dev.json
|
| 35 |
+
visual_features:
|
| 36 |
+
sign2vec:
|
| 37 |
+
enable_input: False
|
| 38 |
+
test: sign2vec/metadata_sign2vec.dev.json
|
| 39 |
+
mae:
|
| 40 |
+
enable_input: False
|
| 41 |
+
test: mae/metadata_mae.dev.json
|
| 42 |
+
dino:
|
| 43 |
+
enable_input: False
|
| 44 |
+
test: dino/metadata_dino.dev.json
|
| 45 |
+
pose:
|
| 46 |
+
enable_input: True
|
| 47 |
+
test: YouTubeASL.keypoints.dev.json
|
| 48 |
+
|
| 49 |
+
SignModelArguments:
|
| 50 |
+
projectors:
|
| 51 |
+
sign2vec:
|
| 52 |
+
dim: 768
|
| 53 |
+
mae:
|
| 54 |
+
dim: 768
|
| 55 |
+
dino:
|
| 56 |
+
dim: 1152
|
| 57 |
+
pose:
|
| 58 |
+
dim: 208
|
predict_pose.py
ADDED
|
@@ -0,0 +1,529 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
from copy import deepcopy
|
| 3 |
+
from typing import List
|
| 4 |
+
|
| 5 |
+
import cv2
|
| 6 |
+
import mediapipe as mp
|
| 7 |
+
import numpy as np
|
| 8 |
+
from mediapipe.tasks import python
|
| 9 |
+
from mediapipe.tasks.python import vision
|
| 10 |
+
from scipy.optimize import linear_sum_assignment
|
| 11 |
+
from ultralytics import YOLO
|
| 12 |
+
|
| 13 |
+
import os
|
| 14 |
+
import json
|
| 15 |
+
import numpy as np
|
| 16 |
+
from datetime import datetime
|
| 17 |
+
import cv2
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def crop_frame(image, bounding_box):
|
| 21 |
+
x, y, w, h = bounding_box
|
| 22 |
+
cropped_frame = image[y:y + h, x:x + w]
|
| 23 |
+
return cropped_frame
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def get_centered_box(keypoints, box_size, scale_factor=1.2):
|
| 27 |
+
center_x, center_y = np.mean(keypoints, axis=0, dtype=int)
|
| 28 |
+
half_size = box_size // 2
|
| 29 |
+
x = center_x - half_size
|
| 30 |
+
y = center_y - half_size
|
| 31 |
+
w = box_size
|
| 32 |
+
h = box_size
|
| 33 |
+
|
| 34 |
+
w_padding = int((scale_factor - 1) * w / 2)
|
| 35 |
+
h_padding = int((scale_factor - 1) * h / 2)
|
| 36 |
+
x -= w_padding
|
| 37 |
+
y -= h_padding
|
| 38 |
+
w += 2 * w_padding
|
| 39 |
+
h += 2 * h_padding
|
| 40 |
+
|
| 41 |
+
return x, y, w, h
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def get_bounding_box(keypoints, scale_factor=1.2):
|
| 45 |
+
keypoints = np.round(keypoints).astype(int)
|
| 46 |
+
x, y, w, h = cv2.boundingRect(keypoints)
|
| 47 |
+
w_padding = int((scale_factor - 1) * w / 2)
|
| 48 |
+
h_padding = int((scale_factor - 1) * h / 2)
|
| 49 |
+
x -= w_padding
|
| 50 |
+
y -= h_padding
|
| 51 |
+
w += 2 * w_padding
|
| 52 |
+
h += 2 * h_padding
|
| 53 |
+
return x, y, w, h
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def adjust_bounding_box(bounding_box, image_shape):
|
| 57 |
+
x, y, w, h = bounding_box
|
| 58 |
+
ih, iw, _ = image_shape
|
| 59 |
+
|
| 60 |
+
# Adjust x-coordinate if the bounding box extends beyond the image's right edge
|
| 61 |
+
if x + w > iw:
|
| 62 |
+
x = iw - w
|
| 63 |
+
|
| 64 |
+
# Adjust y-coordinate if the bounding box extends beyond the image's bottom edge
|
| 65 |
+
if y + h > ih:
|
| 66 |
+
y = ih - h
|
| 67 |
+
|
| 68 |
+
# Ensure bounding box's x and y coordinates are not negative
|
| 69 |
+
x = max(x, 0)
|
| 70 |
+
y = max(y, 0)
|
| 71 |
+
|
| 72 |
+
return x, y, w, h
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
def create_mediapipe_models(checkpoint_folder: str, min_confidence: float = 0.4) -> (object, object, object, object):
|
| 76 |
+
BaseOptions = mp.tasks.BaseOptions
|
| 77 |
+
|
| 78 |
+
# mediapipe
|
| 79 |
+
num_poses = 1
|
| 80 |
+
hand_model_path = os.path.join(checkpoint_folder, 'hand_landmarker.task')
|
| 81 |
+
pose_model_path = os.path.join(checkpoint_folder, 'pose_landmarker_full.task')
|
| 82 |
+
face_model_path = os.path.join(checkpoint_folder, 'face_landmarker.task')
|
| 83 |
+
yolo_model_path = os.path.join(checkpoint_folder, "yolov8n-pose.pt")
|
| 84 |
+
|
| 85 |
+
# yolov8
|
| 86 |
+
yolo_model = YOLO(yolo_model_path)
|
| 87 |
+
|
| 88 |
+
# define hand model
|
| 89 |
+
hand_options = vision.HandLandmarkerOptions(
|
| 90 |
+
base_options=BaseOptions(model_asset_path=hand_model_path),
|
| 91 |
+
min_hand_detection_confidence=min_confidence,
|
| 92 |
+
min_hand_presence_confidence=min_confidence,
|
| 93 |
+
num_hands=num_poses * 2)
|
| 94 |
+
hand_detector = vision.HandLandmarker.create_from_options(hand_options)
|
| 95 |
+
|
| 96 |
+
# define body model
|
| 97 |
+
pose_options = vision.PoseLandmarkerOptions(
|
| 98 |
+
base_options=BaseOptions(model_asset_path=pose_model_path),
|
| 99 |
+
min_pose_detection_confidence=min_confidence,
|
| 100 |
+
min_pose_presence_confidence=min_confidence,
|
| 101 |
+
num_poses=num_poses
|
| 102 |
+
)
|
| 103 |
+
pose_detector = vision.PoseLandmarker.create_from_options(pose_options)
|
| 104 |
+
|
| 105 |
+
# define face model
|
| 106 |
+
face_options = vision.FaceLandmarkerOptions(
|
| 107 |
+
base_options=BaseOptions(model_asset_path=face_model_path),
|
| 108 |
+
min_face_detection_confidence=min_confidence,
|
| 109 |
+
min_face_presence_confidence=min_confidence,
|
| 110 |
+
num_faces=num_poses
|
| 111 |
+
)
|
| 112 |
+
face_detector = vision.FaceLandmarker.create_from_options(face_options)
|
| 113 |
+
|
| 114 |
+
return hand_detector, pose_detector, face_detector, yolo_model
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
def yolo_predict(image: np.ndarray, model, min_conf: float = 0.5):
|
| 118 |
+
yolo_results = model(image, verbose=False)
|
| 119 |
+
|
| 120 |
+
bboxes = yolo_results[0].boxes.xyxy
|
| 121 |
+
keypoints = yolo_results[0].keypoints.xy
|
| 122 |
+
bboxes = bboxes.cpu().numpy()
|
| 123 |
+
keypoints = keypoints.cpu().numpy()
|
| 124 |
+
|
| 125 |
+
conf = yolo_results[0].boxes.conf
|
| 126 |
+
conf = conf.cpu().numpy()
|
| 127 |
+
select_mask_kp = np.sum(keypoints, axis=(1, 2)) > 0.0001
|
| 128 |
+
select_mask_bb = conf > min_conf
|
| 129 |
+
select_mask = select_mask_kp & select_mask_bb
|
| 130 |
+
|
| 131 |
+
conf = conf[select_mask]
|
| 132 |
+
bboxes = bboxes[select_mask]
|
| 133 |
+
keypoints = keypoints[select_mask]
|
| 134 |
+
|
| 135 |
+
return bboxes, keypoints, conf
|
| 136 |
+
|
| 137 |
+
|
| 138 |
+
def load_video_cv(path: str):
|
| 139 |
+
video = []
|
| 140 |
+
|
| 141 |
+
cap = cv2.VideoCapture(path)
|
| 142 |
+
fps = cap.get(cv2.CAP_PROP_FPS)
|
| 143 |
+
ret = True
|
| 144 |
+
while ret:
|
| 145 |
+
ret, img = cap.read()
|
| 146 |
+
if ret:
|
| 147 |
+
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
| 148 |
+
video.append(img)
|
| 149 |
+
cap.release()
|
| 150 |
+
return video, fps
|
| 151 |
+
|
| 152 |
+
|
| 153 |
+
def new_bbox(image, keypoints, lsi=5, rsi=6, sign_space=5):
|
| 154 |
+
h, w = image.shape[:2]
|
| 155 |
+
l_shoulder = keypoints[lsi]
|
| 156 |
+
r_shoulder = keypoints[rsi]
|
| 157 |
+
distance = np.sqrt((l_shoulder[0] - r_shoulder[0]) ** 2 + (l_shoulder[1] - r_shoulder[1]) ** 2)
|
| 158 |
+
|
| 159 |
+
center_x = np.abs(l_shoulder[0] - r_shoulder[0]) / 2 + np.min([l_shoulder[0], r_shoulder[0]], 0)
|
| 160 |
+
center_y = np.abs(l_shoulder[1] - r_shoulder[1]) / 2 + np.min([l_shoulder[1], r_shoulder[1]], 0)
|
| 161 |
+
|
| 162 |
+
new_x0 = center_x - (distance * (sign_space / 2))
|
| 163 |
+
new_x1 = center_x + (distance * (sign_space / 2))
|
| 164 |
+
new_y0 = center_y - (distance * (sign_space / 2))
|
| 165 |
+
new_y1 = center_y + (distance * (sign_space / 2))
|
| 166 |
+
|
| 167 |
+
idx_x = keypoints[:, 0] > 0
|
| 168 |
+
idx_y = keypoints[:, 1] > 0
|
| 169 |
+
new_x0 = np.min([new_x0, *keypoints[idx_x, 0]])
|
| 170 |
+
new_x1 = np.max([new_x1, *keypoints[idx_x, 0]])
|
| 171 |
+
new_y0 = np.min([new_y0, *keypoints[idx_y, 1]])
|
| 172 |
+
new_y1 = np.max([new_y1, *keypoints[idx_y, 1]])
|
| 173 |
+
|
| 174 |
+
new_x0 = np.round(np.clip(new_x0, 0, w)).astype(int)
|
| 175 |
+
new_x1 = np.round(np.clip(new_x1, 0, w)).astype(int)
|
| 176 |
+
new_y0 = np.round(np.clip(new_y0, 0, h)).astype(int)
|
| 177 |
+
new_y1 = np.round(np.clip(new_y1, 0, h)).astype(int)
|
| 178 |
+
|
| 179 |
+
return new_x0, new_y0, new_x1, new_y1
|
| 180 |
+
|
| 181 |
+
|
| 182 |
+
def mdeiapipe_to_xy(data, image_size=None):
|
| 183 |
+
"""image_size: (height, width)"""
|
| 184 |
+
x = np.array([kp.x for kp in data])
|
| 185 |
+
y = np.array([kp.y for kp in data])
|
| 186 |
+
|
| 187 |
+
if image_size is not None:
|
| 188 |
+
x = x * image_size[1]
|
| 189 |
+
y = y * image_size[0]
|
| 190 |
+
|
| 191 |
+
return x, y
|
| 192 |
+
|
| 193 |
+
|
| 194 |
+
def crop_pad_image(image: np.ndarray, bbox: np.ndarray, border: float = 0.25) -> np.ndarray:
|
| 195 |
+
"""Crop the image, pad to square and add a border."""
|
| 196 |
+
# get bbox and image
|
| 197 |
+
x0, y0, x1, y1 = bbox
|
| 198 |
+
w, h = x1 - x0, y1 - y0
|
| 199 |
+
|
| 200 |
+
# add padding
|
| 201 |
+
dif = np.abs(w - h)
|
| 202 |
+
pad_value_0 = np.floor(dif / 2).astype(int)
|
| 203 |
+
pad_value_1 = dif - pad_value_0
|
| 204 |
+
|
| 205 |
+
if w > h:
|
| 206 |
+
y0 -= pad_value_0
|
| 207 |
+
y1 += pad_value_1
|
| 208 |
+
else:
|
| 209 |
+
x0 -= pad_value_0
|
| 210 |
+
x1 += pad_value_1
|
| 211 |
+
|
| 212 |
+
border = np.round((np.max([w, h]) * border) / 2).astype(int)
|
| 213 |
+
ih, iw = image.shape[:2]
|
| 214 |
+
y0 -= border
|
| 215 |
+
y1 += border
|
| 216 |
+
x0 -= border
|
| 217 |
+
x1 += border
|
| 218 |
+
|
| 219 |
+
new_bbox = [x0, y0, x1, y1]
|
| 220 |
+
|
| 221 |
+
y0 += ih
|
| 222 |
+
y1 += ih
|
| 223 |
+
x0 += iw
|
| 224 |
+
x1 += iw
|
| 225 |
+
|
| 226 |
+
image = np.pad(image, ((ih, ih), (iw, iw), (0, 0)), mode='constant', constant_values=0) # mode="reflect"
|
| 227 |
+
cropped_image = image[y0:y1, x0:x1]
|
| 228 |
+
|
| 229 |
+
return cropped_image, new_bbox
|
| 230 |
+
|
| 231 |
+
|
| 232 |
+
def keypoints_out_format(mp_keypoints, image_size):
|
| 233 |
+
"""image_size = (ih, iw)"""
|
| 234 |
+
if len(mp_keypoints) >= 1:
|
| 235 |
+
data = mp_keypoints[0]
|
| 236 |
+
x, y = mdeiapipe_to_xy(data, image_size)
|
| 237 |
+
z = np.array([kp.z for kp in data])
|
| 238 |
+
visibility = np.array([kp.visibility for kp in data])
|
| 239 |
+
data = np.array([x, y, z, visibility]).T
|
| 240 |
+
return data
|
| 241 |
+
else:
|
| 242 |
+
return []
|
| 243 |
+
|
| 244 |
+
|
| 245 |
+
def distance_matrix(P, Q):
|
| 246 |
+
dis_max = np.zeros([len(P), len(Q)])
|
| 247 |
+
for i, p in enumerate(P):
|
| 248 |
+
for j, q in enumerate(Q):
|
| 249 |
+
dist = np.linalg.norm(np.array(p) - np.array(q))
|
| 250 |
+
dis_max[i, j] = dist
|
| 251 |
+
return dis_max
|
| 252 |
+
|
| 253 |
+
|
| 254 |
+
def process_hands(mp_hand_keypoints, mp_handedness, pose_keypoints, image_size, yolo_pose_keypoints=None):
|
| 255 |
+
out = {"left": [], "right": []}
|
| 256 |
+
|
| 257 |
+
if len(mp_hand_keypoints) == 0:
|
| 258 |
+
return out
|
| 259 |
+
|
| 260 |
+
# transform keypoints
|
| 261 |
+
hand_keypoints = []
|
| 262 |
+
for data in mp_hand_keypoints:
|
| 263 |
+
hand_keypoints.append(keypoints_out_format([data], image_size))
|
| 264 |
+
|
| 265 |
+
if (mp_hand_keypoints) == 1:
|
| 266 |
+
side = mp_handedness[0]["category_name"].lower
|
| 267 |
+
out[side] = hand_keypoints[0]
|
| 268 |
+
return out
|
| 269 |
+
|
| 270 |
+
# calculate centers
|
| 271 |
+
hand_centers = []
|
| 272 |
+
for keypoints in hand_keypoints:
|
| 273 |
+
x = keypoints[0, 0]
|
| 274 |
+
y = keypoints[0, 1]
|
| 275 |
+
hand_center = [x, y]
|
| 276 |
+
hand_centers.append(hand_center)
|
| 277 |
+
|
| 278 |
+
# assign hands to sides
|
| 279 |
+
left_wrist = None
|
| 280 |
+
right_wrist = None
|
| 281 |
+
|
| 282 |
+
pose_keypoints = None if len(pose_keypoints) == 0 else pose_keypoints
|
| 283 |
+
if pose_keypoints is not None:
|
| 284 |
+
left_wrist = pose_keypoints[15, :2]
|
| 285 |
+
right_wrist = pose_keypoints[16, :2]
|
| 286 |
+
|
| 287 |
+
elif pose_keypoints is None and yolo_pose_keypoints is not None:
|
| 288 |
+
left_wrist = yolo_pose_keypoints[9, :2]
|
| 289 |
+
right_wrist = yolo_pose_keypoints[10, :2]
|
| 290 |
+
if (np.sum(left_wrist) == 0) or (np.sum(right_wrist) == 0):
|
| 291 |
+
left_wrist = None
|
| 292 |
+
right_wrist = None
|
| 293 |
+
|
| 294 |
+
if left_wrist is not None and right_wrist is not None:
|
| 295 |
+
wrists = [left_wrist, right_wrist]
|
| 296 |
+
|
| 297 |
+
dis_max = distance_matrix(wrists, hand_centers)
|
| 298 |
+
row_idx, col_idx = linear_sum_assignment(dis_max)
|
| 299 |
+
|
| 300 |
+
sides = list(out.keys())
|
| 301 |
+
for ridx, cidx in zip(row_idx, col_idx):
|
| 302 |
+
side = sides[ridx]
|
| 303 |
+
keypoints = hand_keypoints[cidx]
|
| 304 |
+
out[side] = keypoints
|
| 305 |
+
else:
|
| 306 |
+
hand_centers_x = np.array(hand_centers)[:, 0]
|
| 307 |
+
right_idx = np.argmin(hand_centers_x)
|
| 308 |
+
out["right"] = hand_keypoints[right_idx]
|
| 309 |
+
left_idx = np.argmax(hand_centers_x)
|
| 310 |
+
if right_idx != left_idx:
|
| 311 |
+
out["left"] = hand_keypoints[left_idx]
|
| 312 |
+
|
| 313 |
+
return out
|
| 314 |
+
|
| 315 |
+
|
| 316 |
+
def predict_pose(video: List[np.ndarray], models: tuple, sign_space=4, yolo_sign_space=4) -> dict:
|
| 317 |
+
"""
|
| 318 |
+
This function processes a video to detect and extract pose, hand, and face landmarks using Mediapipe models.
|
| 319 |
+
It also calculates the signing space and crops the images accordingly.
|
| 320 |
+
|
| 321 |
+
Parameters:
|
| 322 |
+
video (list): A list of images.
|
| 323 |
+
models (tuple): A tuple containing the Mediapipe models for pose, hand, and face detection and yolo model.
|
| 324 |
+
sign_space (int): The desired size of the signing space.
|
| 325 |
+
Width and height calculated as shoulder distance * sign_space Default is 4.
|
| 326 |
+
|
| 327 |
+
Returns:
|
| 328 |
+
(dict): A dictionary containing the processed video data, including images, keypoints, cropped images, cropped keypoints,
|
| 329 |
+
signing space, and bounding boxes for different body parts.
|
| 330 |
+
"""
|
| 331 |
+
hand_detector, pose_detector, face_detector, yolo_model = models
|
| 332 |
+
results = {
|
| 333 |
+
"images": video,
|
| 334 |
+
"keypoints": [],
|
| 335 |
+
"cropped_images": [],
|
| 336 |
+
"cropped_keypoints": [],
|
| 337 |
+
"sign_space": [],
|
| 338 |
+
"cropped_left_hand": [],
|
| 339 |
+
"cropped_right_hand": [],
|
| 340 |
+
"cropped_face": [],
|
| 341 |
+
"bbox_left_hand": [],
|
| 342 |
+
"bbox_right_hand": [],
|
| 343 |
+
"bbox_face": [],
|
| 344 |
+
}
|
| 345 |
+
|
| 346 |
+
# yolo predict + crop images
|
| 347 |
+
yolo_predictions = []
|
| 348 |
+
num_predictions = []
|
| 349 |
+
for idx, image in enumerate(results["images"]):
|
| 350 |
+
bboxes, keypoints, confs = yolo_predict(image, yolo_model)
|
| 351 |
+
yolo_predictions.append([bboxes, keypoints, confs])
|
| 352 |
+
num_predictions.append(len(bboxes))
|
| 353 |
+
|
| 354 |
+
# no predictions -> add empty values and return
|
| 355 |
+
if np.sum(num_predictions) == 0:
|
| 356 |
+
_h, _w = results["images"][0].shape[:2]
|
| 357 |
+
for idx in range(len(results["images"])):
|
| 358 |
+
results["keypoints"].append({'pose_landmarks': [], 'right_hand_landmarks': [], 'left_hand_landmarks': [], 'face_landmarks': []})
|
| 359 |
+
results["cropped_images"].append(results["images"][idx])
|
| 360 |
+
results["cropped_keypoints"].append({'pose_landmarks': [], 'right_hand_landmarks': [], 'left_hand_landmarks': [], 'face_landmarks': []})
|
| 361 |
+
results["sign_space"].append([0, 0, _w, _h])
|
| 362 |
+
results["cropped_left_hand"].append(np.zeros([224, 224, 3], dtype=np.uint8))
|
| 363 |
+
results["cropped_right_hand"].append(np.zeros([224, 224, 3], dtype=np.uint8))
|
| 364 |
+
results["cropped_face"].append(np.zeros([224, 224, 3], dtype=np.uint8))
|
| 365 |
+
results["bbox_left_hand"].append([])
|
| 366 |
+
results["bbox_right_hand"].append([])
|
| 367 |
+
results["bbox_face"].append([])
|
| 368 |
+
return results
|
| 369 |
+
|
| 370 |
+
# get signing bbox
|
| 371 |
+
x0, y0, x1, y1 = [], [], [], []
|
| 372 |
+
for idx, (image, prediction) in enumerate(zip(results["images"], yolo_predictions)):
|
| 373 |
+
_, keypoints, _ = prediction
|
| 374 |
+
if len(keypoints) != 1:
|
| 375 |
+
continue
|
| 376 |
+
|
| 377 |
+
_x0, _y0, _x1, _y1 = new_bbox(image, keypoints[0], lsi=5, rsi=6, sign_space=yolo_sign_space)
|
| 378 |
+
|
| 379 |
+
x0.append(_x0)
|
| 380 |
+
y0.append(_y0)
|
| 381 |
+
x1.append(_x1)
|
| 382 |
+
y1.append(_y1)
|
| 383 |
+
|
| 384 |
+
x0y = np.round(np.median(x0)).astype(int)
|
| 385 |
+
y0y = np.round(np.median(y0)).astype(int)
|
| 386 |
+
x1y = np.round(np.median(x1)).astype(int)
|
| 387 |
+
y1y = np.round(np.median(y1)).astype(int)
|
| 388 |
+
|
| 389 |
+
# mediapipe predict + signing space
|
| 390 |
+
mp_predictions = []
|
| 391 |
+
x0, y0, x1, y1 = [], [], [], []
|
| 392 |
+
for idx, image in enumerate(results["images"]):
|
| 393 |
+
yolo_image = image[y0y:y1y, x0y:x1y]
|
| 394 |
+
|
| 395 |
+
ih, iw = yolo_image.shape[:2]
|
| 396 |
+
mp_image = mp.Image(image_format=mp.ImageFormat.SRGB, data=np.array(yolo_image))
|
| 397 |
+
|
| 398 |
+
# HACK:
|
| 399 |
+
# if the YOLO model does not detect anything,
|
| 400 |
+
# pretend it detects a black square and give it to mediapipe
|
| 401 |
+
if yolo_image.shape == (0, 0, 3):
|
| 402 |
+
mp_image = mp.Image(
|
| 403 |
+
image_format=mp.ImageFormat.SRGB,
|
| 404 |
+
data=np.zeros(shape=(256, 256, 3), dtype=np.uint8)
|
| 405 |
+
)
|
| 406 |
+
|
| 407 |
+
pose_prediction = pose_detector.detect(mp_image)
|
| 408 |
+
hand_prediction = hand_detector.detect(mp_image)
|
| 409 |
+
face_prediction = face_detector.detect(mp_image)
|
| 410 |
+
|
| 411 |
+
mp_predictions.append([hand_prediction, face_prediction, pose_prediction])
|
| 412 |
+
|
| 413 |
+
if len(pose_prediction.pose_landmarks) != 1:
|
| 414 |
+
continue
|
| 415 |
+
|
| 416 |
+
kp_all_x = []
|
| 417 |
+
kp_all_y = []
|
| 418 |
+
mp_keypoints = [
|
| 419 |
+
pose_prediction.pose_landmarks[0][:25],
|
| 420 |
+
*face_prediction.face_landmarks,
|
| 421 |
+
*hand_prediction.hand_landmarks,
|
| 422 |
+
]
|
| 423 |
+
|
| 424 |
+
for p in mp_keypoints:
|
| 425 |
+
x, y = mdeiapipe_to_xy(p, (ih, iw))
|
| 426 |
+
kp_all_x.extend(x)
|
| 427 |
+
kp_all_y.extend(y)
|
| 428 |
+
kp_all = np.array((kp_all_x, kp_all_y)).T
|
| 429 |
+
|
| 430 |
+
kp_all[:, 0] = kp_all[:, 0] + x0y
|
| 431 |
+
kp_all[:, 1] = kp_all[:, 1] + y0y
|
| 432 |
+
|
| 433 |
+
if len(kp_all) == 0:
|
| 434 |
+
continue
|
| 435 |
+
|
| 436 |
+
_x0, _y0, _x1, _y1 = new_bbox(image, kp_all, lsi=11, rsi=12, sign_space=sign_space)
|
| 437 |
+
|
| 438 |
+
x0.append(_x0)
|
| 439 |
+
y0.append(_y0)
|
| 440 |
+
x1.append(_x1)
|
| 441 |
+
y1.append(_y1)
|
| 442 |
+
|
| 443 |
+
# create signing space as median of all signing spaces
|
| 444 |
+
if len(x0) == 0:
|
| 445 |
+
ih, iw = video[0].shape[:2]
|
| 446 |
+
x0mp = 0
|
| 447 |
+
y0mp = 0
|
| 448 |
+
x1mp = iw
|
| 449 |
+
y1mp = ih
|
| 450 |
+
else:
|
| 451 |
+
x0mp = np.round(np.median(x0)).astype(int)
|
| 452 |
+
y0mp = np.round(np.median(y0)).astype(int)
|
| 453 |
+
x1mp = np.round(np.median(x1)).astype(int)
|
| 454 |
+
y1mp = np.round(np.median(y1)).astype(int)
|
| 455 |
+
|
| 456 |
+
for idx, (image, prediction) in enumerate(zip(results["images"], mp_predictions)):
|
| 457 |
+
yolo_image = image[y0y:y1y, x0y:x1y]
|
| 458 |
+
yih, yiw = yolo_image.shape[:2]
|
| 459 |
+
|
| 460 |
+
cropped_image, pad_bbox = crop_pad_image(image, (x0mp, y0mp, x1mp, y1mp), border=0)
|
| 461 |
+
|
| 462 |
+
hand_prediction, face_prediction, pose_prediction = prediction
|
| 463 |
+
|
| 464 |
+
face_keypoints = keypoints_out_format(face_prediction.face_landmarks, (yih, yiw))
|
| 465 |
+
pose_keypoints = keypoints_out_format(pose_prediction.pose_landmarks, (yih, yiw))
|
| 466 |
+
hand_keypoints = process_hands(
|
| 467 |
+
hand_prediction.hand_landmarks,
|
| 468 |
+
hand_prediction.handedness,
|
| 469 |
+
pose_keypoints,
|
| 470 |
+
(yih, yiw),
|
| 471 |
+
None
|
| 472 |
+
)
|
| 473 |
+
|
| 474 |
+
keypoints = {
|
| 475 |
+
'pose_landmarks': pose_keypoints,
|
| 476 |
+
'right_hand_landmarks': hand_keypoints["right"],
|
| 477 |
+
'left_hand_landmarks': hand_keypoints["left"],
|
| 478 |
+
'face_landmarks': face_keypoints
|
| 479 |
+
}
|
| 480 |
+
|
| 481 |
+
# move kp
|
| 482 |
+
x_move = x0y
|
| 483 |
+
y_move = y0y
|
| 484 |
+
for name in keypoints:
|
| 485 |
+
if len(keypoints[name]) > 0:
|
| 486 |
+
keypoints[name][:, 0] += x_move
|
| 487 |
+
keypoints[name][:, 1] += y_move
|
| 488 |
+
|
| 489 |
+
# get dino crops
|
| 490 |
+
name_to_keypoints = [
|
| 491 |
+
("face", face_keypoints),
|
| 492 |
+
("left_hand", hand_keypoints["left"]),
|
| 493 |
+
("right_hand", hand_keypoints["right"])
|
| 494 |
+
]
|
| 495 |
+
for name, kp in name_to_keypoints:
|
| 496 |
+
if len(kp) > 0:
|
| 497 |
+
kp = np.round(kp[:, :2]).astype(int)
|
| 498 |
+
x, y, w, h = cv2.boundingRect(kp)
|
| 499 |
+
cropped_local_bbox = get_centered_box(kp, np.max([w, h]), scale_factor=1.2)
|
| 500 |
+
cropped_local_bbox = adjust_bounding_box(cropped_local_bbox, image.shape)
|
| 501 |
+
cropped_local_image = crop_frame(image, cropped_local_bbox)
|
| 502 |
+
x0, y0, w, h = cropped_local_bbox
|
| 503 |
+
cropped_local_bbox = [x0, y0, x0 + w, y0 + h]
|
| 504 |
+
|
| 505 |
+
else:
|
| 506 |
+
cropped_local_image = np.zeros([224, 224, 3], dtype=np.uint8)
|
| 507 |
+
cropped_local_bbox = []
|
| 508 |
+
results[f"bbox_{name}"].append(cropped_local_bbox)
|
| 509 |
+
results[f"cropped_{name}"].append(cropped_local_image)
|
| 510 |
+
|
| 511 |
+
# move kp
|
| 512 |
+
x_move = pad_bbox[0]
|
| 513 |
+
y_move = pad_bbox[1]
|
| 514 |
+
keypoints_cropped = deepcopy(keypoints)
|
| 515 |
+
for name in keypoints_cropped:
|
| 516 |
+
if len(keypoints_cropped[name]) > 0:
|
| 517 |
+
keypoints_cropped[name][:, 0] -= x_move
|
| 518 |
+
keypoints_cropped[name][:, 1] -= y_move
|
| 519 |
+
# Slice [:, :2] to keep only the first two columns before saving
|
| 520 |
+
keypoints_cropped[name] = np.round(keypoints_cropped[name][:, :2], 3).tolist()
|
| 521 |
+
keypoints[name] = np.round(keypoints[name][:, :2], 3).tolist()
|
| 522 |
+
# save processed data
|
| 523 |
+
results["keypoints"].append(keypoints)
|
| 524 |
+
results["cropped_images"].append(cropped_image)
|
| 525 |
+
results["cropped_keypoints"].append(keypoints_cropped)
|
| 526 |
+
results["sign_space"].append(pad_bbox)
|
| 527 |
+
results["images"] = video
|
| 528 |
+
|
| 529 |
+
return results
|
utils/__pycache__/translation.cpython-311.pyc
ADDED
|
Binary file (2.93 kB). View file
|
|
|
utils/translation.py
ADDED
|
@@ -0,0 +1,32 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
|
| 3 |
+
def postprocess_text(preds, labels):
|
| 4 |
+
preds = [pred.strip() for pred in preds]
|
| 5 |
+
labels = [[label.strip()] for label in labels]
|
| 6 |
+
|
| 7 |
+
return preds, labels
|
| 8 |
+
|
| 9 |
+
# Add collate_fn to DataLoader
|
| 10 |
+
def collate_fn(batch):
|
| 11 |
+
# Add padding to the inputs
|
| 12 |
+
# "inputs" must be 250 frames long
|
| 13 |
+
# "attention_mask" must be 250 frames long
|
| 14 |
+
# "labels" must be 128 tokens long
|
| 15 |
+
return {
|
| 16 |
+
"sign_inputs": torch.stack([
|
| 17 |
+
torch.cat((sample["sign_inputs"], torch.zeros(250 - sample["sign_inputs"].shape[0], 208)), dim=0)
|
| 18 |
+
for sample in batch
|
| 19 |
+
]),
|
| 20 |
+
"attention_mask": torch.stack([
|
| 21 |
+
torch.cat((sample["attention_mask"], torch.zeros(250 - sample["attention_mask"].shape[0])), dim=0)
|
| 22 |
+
if sample["attention_mask"].shape[0] < 250
|
| 23 |
+
else sample["attention_mask"]
|
| 24 |
+
for sample in batch
|
| 25 |
+
]),
|
| 26 |
+
"labels": torch.stack([
|
| 27 |
+
torch.cat((sample["labels"].squeeze(0), torch.zeros(128 - sample["labels"].shape[0])), dim=0)
|
| 28 |
+
if sample["labels"].shape[0] < 128
|
| 29 |
+
else sample["labels"]
|
| 30 |
+
for sample in batch
|
| 31 |
+
]).squeeze(0).to(torch.long),
|
| 32 |
+
}
|