import os import time import torch import numpy as np import matplotlib.pyplot as plt from dotenv import load_dotenv # Import tvého nového modelu a pomocných funkcí from Uni_Sign.models import Uni_Sign from Uni_Sign.datasets import load_part_kp_YTASL, YTASL_GROUP_SIZES, _fill_missing_landmarks, select_frame_indices from predict_pose import create_mediapipe_models, predict_pose, load_video_cv os.environ['KMP_DUPLICATE_LIB_OK'] = 'TRUE' load_dotenv() # Globální proměnné pro modely, ať se nenačítají při každém videu znovu model = None pose_models = None args = None class InferenceConfig: """Čistá konfigurace pro Uni_Sign inferenci""" def __init__(self): # Nezapomeň upravit cestu ke svým natrénovaným váhám self.finetune = os.environ.get("UNISIGN_WEIGHTS", r"./Uni_Sign/unisign_model/best_checkpoint.pth") self.dataset = "YTASL" self.task = "SLT" self.max_length = 256 self.normalization = "none" self.layout = "pruned" self.n_registers = 0 self.hidden_dim = 256 self.rgb_support = False self.no_adaptive_gcn = False self.register_position = "before_all" self.label_smoothing = 0.2 self.batch_size = 1 def initialize_model(): """Inicializuje Uni_Sign a Mediapipe/YOLO modely pro extrakci pose.""" global model, pose_models, args if model is not None and pose_models is not None: return print("Inicializuji Uni_Sign model...") args = InferenceConfig() model = Uni_Sign(args=args) # Načtení vah if args.finetune and os.path.exists(args.finetune): print(f"Načítám váhy z: {args.finetune}") state_dict = torch.load(args.finetune, map_location='cpu')['model'] model.load_state_dict(state_dict, strict=False) else: print(f"VAROVÁNÍ: Checkpoint '{args.finetune}' nebyl nalezen!") device = "cuda" if torch.cuda.is_available() else "cpu" model.to(device) model.eval() print("Inicializuji Mediapipe a YOLO modely...") pose_checkpoint_folder = 'checkpoints/pose/' pose_models = create_mediapipe_models(pose_checkpoint_folder) def process_pose_data_in_memory(pose_results, args): """ Nahrazuje starý 'process_single_json'. Zpracovává 'cropped_keypoints' přímo ze slovníku (v paměti). """ raw_pose = pose_results.get('cropped_keypoints', []) if not raw_pose: raise ValueError("Video neobsahuje žádné detekované pose (cropped_keypoints).") # NOVÝ KÓD: Převod numpy polí na čisté [X, Y] seznamy a ošetření chybějících částí těla pose = [] for frame_data in raw_pose: formatted_frame = {} for part, expected_size in YTASL_GROUP_SIZES.items(): kps = frame_data.get(part, []) # Pokud část chybí (prázdný list/array), pošleme prázdno, ať si s tím poradí _fill_missing_landmarks # Pokud část chybí (prázdný list/array), pošleme prázdno # Pokud část chybí (prázdný list/array), pošleme prázdno # Bezpečná kontrola Numpy pole: if kps is None or len(kps) == 0: formatted_frame[part] = [] else: # Neprůstřelné řešení: převedeme data (ať už jsou cokoliv) na Numpy pole, # ořízneme první dva sloupce a převedeme zpět na čistý list. formatted_frame[part] = np.array(kps)[:, :2].tolist() pose.append(formatted_frame) # ... Zbytek funkce pokračuje beze změny: duration = len(pose) tmp = select_frame_indices(duration, args.max_length, phase='test') skeletons = [pose[i] for i in tmp] confs = [] for i, skeleton in enumerate(skeletons): conf = {} for group_name, expected_size in YTASL_GROUP_SIZES.items(): _fill_missing_landmarks( skeleton=skeleton, conf=conf, group_name=group_name, expected_size=expected_size, clip_name="gradio_video", frame_idx=i, error_group_label=f"group '{group_name}'", include_size_details=False, strict_key_access=True, ) confs.append(conf) kps_with_scores = load_part_kp_YTASL(skeletons, confs, args.normalization, args.layout) src_input = {} for key, val in kps_with_scores.items(): src_input[key] = val.unsqueeze(0) seq_len = src_input['body'].shape[1] mask_gen = torch.ones([seq_len]) + 7 src_input['attention_mask'] = (mask_gen != 0).long().unsqueeze(0) src_input['src_length_batch'] = torch.LongTensor([seq_len]) src_input['name_batch'] = ["gradio_video"] return src_input def process_input(input_video_path): """ Hlavní vstupní bod pro Gradio. Spustí extrakci dat, vizualizaci a inference přes Uni_Sign. """ try: initialize_model() print("1. Extrahuji keypoints z videa...") video_frames, _ = load_video_cv(input_video_path) pose_results = predict_pose(video_frames, pose_models) print("2. Prostor pro vizualizaci...") ''' print("3. Přenáším data do Tenzorů pro Uni_Sign...") src_input = process_pose_data_in_memory(pose_results, args) # Přesunutí Tenzorů na stejné zařízení (GPU/CPU) jako model device = next(model.parameters()).device for key in ['body', 'left', 'right', 'face_all', 'attention_mask']: if key in src_input: # Ošetření typu dat na float if key != 'attention_mask': src_input[key] = src_input[key].to(device, dtype=torch.float32) else: src_input[key] = src_input[key].to(device) tgt_input = {'gt_sentence': [""], 'gt_gloss': [""]} print("4. Generuji překlad...") with torch.no_grad(): stack_out = model(src_input, tgt_input) output = model.generate( stack_out, max_new_tokens=100, num_beams=4, ) tokenizer = model.mt5_tokenizer tgt_pres = tokenizer.batch_decode(output, skip_special_tokens=True) result = tgt_pres[0].strip() if not result: return "Model nedokázal generovat text. Zkus jiné video." print("5. VŠE HOTOVO!") return result ''' print("3. Přenáším data do Tenzorů pro Uni_Sign...") try: print("DEBUG 3: Spouštím 'process_pose_data_in_memory' (převod do PyTorch)...") src_input = process_pose_data_in_memory(pose_results, args) print(f"DEBUG 3: Funkce úspěšně dokončena. Nalezené klíče: {list(src_input.keys())}") # Přesunutí Tenzorů na stejné zařízení (GPU/CPU) jako model device = next(model.parameters()).device print(f"DEBUG 3: Zařízení modelu je nastaveno na: {device}. Přesouvám tensory...") for key in ['body', 'left', 'right', 'face_all', 'attention_mask']: if key in src_input: puvodni_tvar = src_input[key].shape # Ošetření typu dat na float if key != 'attention_mask': src_input[key] = src_input[key].to(device, dtype=torch.float32) print(f"DEBUG 3: -- '{key}' (float32) přesunut na {device}. Tvar: {puvodni_tvar}") else: src_input[key] = src_input[key].to(device) print(f"DEBUG 3: -- '{key}' (maska) přesunuta na {device}. Tvar: {puvodni_tvar}") else: print(f"DEBUG 3: KRITICKÉ VAROVÁNÍ - Klíč '{key}' ve výstupu úplně chybí!") except Exception as e: print("=====================================") print("CHYBA V KROKU 3 (PŘÍPRAVA DAT/TENZORŮ)!") print(f"Typ chyby: {type(e).__name__}") print(f"Zpráva chyby: {str(e)}") import traceback traceback.print_exc() print("=====================================") return f"Došlo k chybě při formátování dat (Krok 3): {str(e)}" tgt_input = {'gt_sentence': [""], 'gt_gloss': [""]} print("4. Generuji překlad...") # VYLEPŠENÁ KONTROLA VSTUPŮ S KONTROLOU HODNOT (Zda to nejsou samé nuly) # VYLEPŠENÁ KONTROLA VSTUPŮ S BEZPEČNÝM VÝPOČTEM print(f"DEBUG: Typ src_input: {type(src_input)}") if isinstance(src_input, dict): for k, v in src_input.items(): if hasattr(v, 'shape') and hasattr(v, 'min'): # OPRAVA: Převedeme tenzor na float() jen pro účely výpočtu průměru v_min = v.min().item() v_max = v.max().item() v_mean = v.float().mean().item() print(f"DEBUG: -- src_input['{k}'] typ: {type(v)}, tvar: {v.shape}, min: {v_min:.4f}, max: {v_max:.4f}, mean: {v_mean:.4f}") else: print(f"DEBUG: -- src_input['{k}'] typ: {type(v)}, tvar: {v.shape if hasattr(v, 'shape') else 'Není tensor'}") elif hasattr(src_input, 'shape'): print(f"DEBUG: Tvar src_input: {src_input.shape}") try: with torch.no_grad(): stack_out = model(src_input, tgt_input) tokenizer = model.mt5_tokenizer # NÁVRAT K ORIGINÁLNÍMU VOLÁNÍ (Odebrány nepodporované argumenty) output = model.generate( stack_out, max_new_tokens=100, num_beams=4, ) print(f"DEBUG: Výstupní tensor má tvar: {output.shape}") print(f"DEBUG: Vygenerované token IDs (čísla): {output.tolist()}") tgt_pres = tokenizer.batch_decode(output, skip_special_tokens=True) print(f"DEBUG: Surový text z tokenizéru: '{tgt_pres}'") surove_tokeny = tokenizer.batch_decode(output, skip_special_tokens=False) print(f"DEBUG: Text VČETNĚ speciálních tokenů: '{surove_tokeny}'") result = tgt_pres[0].strip() if not result: return "Model nevygeneroval žádný překlad. Pravděpodobně dostal prázdná nebo špatně zformátovaná vstupní data." print(f"5. VŠE HOTOVO! Překlad: {result}") return result except Exception as e: print("=====================================") print("CHYBA PŘI GENEROVÁNÍ PŘEKLADU!") print(f"Typ chyby: {type(e).__name__}") print(f"Zpráva chyby: {str(e)}") import traceback traceback.print_exc() print("=====================================") return f"Došlo k vnitřní chybě modelu při generování: {str(e)}" except Exception as e: print(f"Chyba při zpracování: {e}") import traceback traceback.print_exc() return f"Chyba při zpracování videa: {str(e)}"