plice13 commited on
Commit
988f575
·
verified ·
1 Parent(s): 91e379c

final demo #1

Browse files
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
- return "DONE", None
 
 
 
 
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
+ }