Spaces:
Running on Zero
Running on Zero
Fix 2p->20s + 20 bugs: pad audio, OCR/ASR coverage gates, TTS cache/placeholder, omni cloud, ASS escape, timeout, preflight, pool locks, SSRF guards
df03342 verified Download app/core/tts_worker_cli.py from hoangtaiii/DRIPPY4: direct link, hf CLI and curl.
- Browser
- Download file 64.5 kB
-
https://huggingface.co/spaces/hoangtaiii/DRIPPY4/resolve/main/app/core/tts_worker_cli.py
- Command line
-
hf download hf://spaces/hoangtaiii/DRIPPY4/app/core/tts_worker_cli.py
-
curl -L -o tts_worker_cli.py https://huggingface.co/spaces/hoangtaiii/DRIPPY4/resolve/main/app/core/tts_worker_cli.py
64.5 kB
| import sys | |
| import os | |
| import re | |
| import json | |
| import argparse | |
| import asyncio | |
| import time | |
| import urllib.request | |
| import urllib.error | |
| import subprocess | |
| import threading | |
| import queue | |
| import uuid | |
| import hashlib | |
| import atexit | |
| from pathlib import Path | |
| # Enforce UTF-8 for Windows console | |
| if sys.platform == 'win32': | |
| try: | |
| if hasattr(sys.stdout, 'reconfigure'): | |
| sys.stdout.reconfigure(encoding='utf-8') | |
| if hasattr(sys.stderr, 'reconfigure'): | |
| sys.stderr.reconfigure(encoding='utf-8') | |
| except Exception: | |
| pass | |
| # Ensure app is in path | |
| sys.path.append(str(Path(__file__).parent.parent.parent)) | |
| from app.core.vietnamese_text_normalizer import VietnameseTextNormalizer | |
| from app.core.audio_timeline_classifier import AudioTimelineClassifier | |
| from app.core.audio_mixer import AudioMixer | |
| from app.core.gpu_resource_manager import GPUResourceManager | |
| def parse_time(time_str): | |
| try: | |
| time_str = time_str.strip().replace('.', ',') | |
| parts = time_str.split(':') | |
| if len(parts) != 3: | |
| raise ValueError(f"Invalid timestamp format: {time_str}") | |
| h, m, s_ms = parts | |
| if ',' not in s_ms: | |
| raise ValueError(f"Invalid timestamp format: {time_str}") | |
| s, ms = s_ms.split(',') | |
| return int(h.strip())*3600000 + int(m.strip())*60000 + int(s.strip())*1000 + int(ms.strip()) | |
| except Exception as e: | |
| print(f"Warning: Failed to parse time string '{time_str}': {e}", file=sys.stderr) | |
| return None | |
| def parse_srt(srt_path): | |
| blocks = [] | |
| if not srt_path or not Path(srt_path).exists(): | |
| return blocks | |
| try: | |
| with open(srt_path, 'r', encoding='utf-8') as f: | |
| content = f.read().strip() | |
| except Exception as e: | |
| print(f"Warning: Failed to read SRT file {srt_path}: {e}", file=sys.stderr) | |
| return blocks | |
| content = content.replace('\r\n', '\n') | |
| pattern = r'(\d+)\s+(\d{2}:\d{2}:\d{2}[.,]\d{3}\s*-->\s*\d{2}:\d{2}:\d{2}[.,]\d{3})\s*\n(.*?)(?=\n\s*\d+\s+\d{2}:\d{2}:\d{2}[.,]\d{3}\s*-->|\Z)' | |
| matches = re.finditer(pattern, content, re.DOTALL) | |
| for match in matches: | |
| try: | |
| b_id = match.group(1).strip() | |
| if b_id.isdigit(): | |
| b_id_int = int(b_id) | |
| else: | |
| continue | |
| timestamp = match.group(2).strip() | |
| times = timestamp.split(' --> ') | |
| if len(times) == 2: | |
| start_ms = parse_time(times[0]) | |
| end_ms = parse_time(times[1]) | |
| if start_ms is None or end_ms is None: | |
| continue | |
| text = match.group(3).strip() | |
| text = " ".join([l.strip() for l in text.split('\n') if l.strip()]) | |
| blocks.append({ | |
| "id": str(b_id_int), | |
| "start_ms": start_ms, | |
| "end_ms": end_ms, | |
| "text": text | |
| }) | |
| except Exception as e: | |
| print(f"Warning: Skipping malformed block: {e}", file=sys.stderr) | |
| continue | |
| return blocks | |
| def verify_audio_file(filepath, min_size=500, min_duration=0.05, ffmpeg_path=None) -> bool: | |
| path = Path(filepath) | |
| if not path.exists(): | |
| return False | |
| if path.stat().st_size < min_size: | |
| return False | |
| try: | |
| from pydub import AudioSegment | |
| if ffmpeg_path: | |
| AudioSegment.converter = str(ffmpeg_path) | |
| audio = AudioSegment.from_file(str(path)) | |
| if len(audio) / 1000.0 < min_duration: | |
| return False | |
| return True | |
| except Exception: | |
| return False | |
| def sanitize_for_edge_tts(text: str) -> str: | |
| replacements = { | |
| "\u2011": "-", # non-breaking hyphen | |
| "\u2010": "-", | |
| "\u2013": "-", | |
| "\u2014": "-", | |
| "\u200b": "", | |
| "\ufeff": "", | |
| } | |
| for src, dst in replacements.items(): | |
| text = text.replace(src, dst) | |
| # Remove emojis and invisible characters using unicode match | |
| text = re.sub(r'[^\w\s,.:;?!@#$\-%&*()\'\"+=–—\/\\’“”]', '', text, flags=re.UNICODE) | |
| text = " ".join(text.split()) | |
| return text.strip() | |
| def split_text_in_half(text: str): | |
| words = text.split() | |
| if len(words) <= 1: | |
| return [text] | |
| mid = len(words) // 2 | |
| best_idx = -1 | |
| for offset in range(mid): | |
| for idx in [mid + offset, mid - offset]: | |
| if 0 <= idx < len(words) - 1: | |
| if words[idx].endswith((',', '.', ';', ':', '?', '!')): | |
| best_idx = idx | |
| break | |
| if best_idx != -1: | |
| break | |
| if best_idx == -1: | |
| best_idx = mid | |
| part1 = " ".join(words[:best_idx+1]) | |
| part2 = " ".join(words[best_idx+1:]) | |
| return [part1, part2] | |
| def write_reports(parent_dir, failed_segments, manifest_entries): | |
| parent_dir = Path(parent_dir) | |
| try: | |
| with open(parent_dir / "failed_tts_segments.json", "w", encoding="utf-8") as f: | |
| json.dump(failed_segments, f, ensure_ascii=False, indent=2) | |
| except Exception as e: | |
| print(f"Warning: Failed to save failed_tts_segments.json: {e}", file=sys.stderr) | |
| try: | |
| with open(parent_dir / "tts_segments_manifest.json", "w", encoding="utf-8") as f: | |
| json.dump(manifest_entries, f, ensure_ascii=False, indent=2) | |
| except Exception as e: | |
| print(f"Warning: Failed to save tts_segments_manifest.json: {e}", file=sys.stderr) | |
| try: | |
| synthesis_failures = [ | |
| row for row in failed_segments | |
| if not str(row.get("reason", "")).startswith("tts_timing_overflow") | |
| ] | |
| timing_overflows = [ | |
| row for row in failed_segments | |
| if str(row.get("reason", "")).startswith("tts_timing_overflow") | |
| ] | |
| report = { | |
| "total_segments": len(manifest_entries), | |
| "failed_segments_count": len(synthesis_failures), | |
| "synthesis_failed_count": len(synthesis_failures), | |
| "timing_overflow_count": len(timing_overflows), | |
| "status": "FAILED" if synthesis_failures else ("NEED_REVIEW" if timing_overflows else "OK"), | |
| "timestamp": time.time() | |
| } | |
| with open(parent_dir / "tts_report.json", "w", encoding="utf-8") as f: | |
| json.dump(report, f, ensure_ascii=False, indent=2) | |
| except Exception as e: | |
| print(f"Warning: Failed to save tts_report.json: {e}", file=sys.stderr) | |
| async def run_edge_tts(text, voice, output_path, speed="1.0", pitch="0", volume="100"): | |
| import edge_tts | |
| # Format edge-tts parameters | |
| rate_str = "+0%" | |
| try: | |
| speed_val = float(speed) | |
| pct = int((speed_val - 1.0) * 100) | |
| rate_str = f"{'+' if pct >= 0 else ''}{pct}%" | |
| except Exception: | |
| rate_str = "+0%" | |
| pitch_str = "+0Hz" | |
| try: | |
| pitch_val = int(pitch) | |
| pitch_str = f"{'+' if pitch_val >= 0 else ''}{pitch_val}Hz" | |
| except Exception: | |
| if "%" in str(pitch): | |
| pitch_str = str(pitch) | |
| volume_str = "+0%" | |
| try: | |
| vol_val = int(volume) | |
| pct = vol_val - 100 | |
| volume_str = f"{'+' if pct >= 0 else ''}{pct}%" | |
| except Exception: | |
| volume_str = "+0%" | |
| communicate = edge_tts.Communicate(text, voice, rate=rate_str, pitch=pitch_str, volume=volume_str) | |
| await communicate.save(str(output_path)) | |
| def download_piper_model(voice_name, dest_dir): | |
| dest_dir = Path(dest_dir) | |
| dest_dir.mkdir(parents=True, exist_ok=True) | |
| onnx_file = dest_dir / f"{voice_name}.onnx" | |
| json_file = dest_dir / f"{voice_name}.onnx.json" | |
| if onnx_file.exists() and json_file.exists(): | |
| return onnx_file, json_file | |
| voice_map = { | |
| "vi_VN-vais1000-medium": "vi/vi_VN/vais1000/medium/vi_VN-vais1000-medium", | |
| "vi_VN-vivos-x_low": "vi/vi_VN/vivos/x_low/vi_VN-vivos-x_low", | |
| "vi_VN-25hours_single-low": "vi/vi_VN/25hours_single/low/vi_VN-25hours_single-low" | |
| } | |
| hf_path = voice_map.get(voice_name, "vi/vi_VN/vais1000/medium/vi_VN-vais1000-medium") | |
| base_url = f"https://huggingface.co/rhasspy/piper-voices/resolve/main/{hf_path}" | |
| print(f"Downloading Piper model files for {voice_name} to {dest_dir}...") | |
| urllib.request.urlretrieve(f"{base_url}.onnx.json", str(json_file)) | |
| urllib.request.urlretrieve(f"{base_url}.onnx", str(onnx_file)) | |
| return onnx_file, json_file | |
| # ── OmniVoice Presets — 100% đồng bộ với D:\omnivoice - web\app.py ── | |
| # Extended mapping: voice_name -> (base_preset, emotion_tag) | |
| # FIX: trước đây 2 hằng này không được định nghĩa/import -> NameError ngay khi | |
| # engine chứa "omnivoice". Kiểu dict (không phải set) theo đúng cách dùng bên dưới. | |
| OMNIVOICE_PRESETS = { | |
| "fun_male": "male, young adult, high pitch", | |
| "story_male": "male, young adult, moderate pitch", | |
| "meme_male": "male, young adult, moderate pitch", | |
| "calm_fun_male": "male, young adult, low pitch", | |
| } | |
| OMNIVOICE_EXTENDED = { | |
| "meme_male_gasp": ("meme_male", "gasp"), | |
| "meme_male_normal": ("meme_male", "normal"), | |
| "meme_male_excited": ("meme_male", "excited"), | |
| "meme_male_confident": ("meme_male", "confident"), | |
| "meme_male_sad": ("meme_male", "sad"), | |
| "meme_male_whispering": ("meme_male", "whispering"), | |
| "fun_male_gasp": ("fun_male", "gasp"), | |
| "fun_male_excited": ("fun_male", "excited"), | |
| "story_male_normal": ("story_male", "normal"), | |
| "calm_fun_male_sad": ("calm_fun_male", "sad"), | |
| "calm_fun_male_whispering": ("calm_fun_male", "whispering"), | |
| } | |
| # Placeholder của translation khi fail — TTS phải SKIP (im lặng) chứ không đọc lên | |
| TTS_PLACEHOLDER_TEXTS = { | |
| "[CẦN DỊCH LẠI]", "[CAN DICH LAI]", "[DỊCH LẠI]", "[CẦN DỊCH]", | |
| "CẦN DỊCH LẠI", "CAN DICH LAI", | |
| } | |
| def _is_tts_placeholder(text: str) -> bool: | |
| t = str(text or "").strip() | |
| if t in TTS_PLACEHOLDER_TEXTS: | |
| return True | |
| # Biến thể khoảng trắng/thường-hoa | |
| norm = re.sub(r"\s+", " ", t).strip().strip("[]").upper() | |
| return norm in {"CẦN DỊCH LẠI", "CAN DICH LAI", "DỊCH LẠI", "CẦN DỊCH"} | |
| def _resolve_omnivoice_preset(voice: str, style: str = ""): | |
| """ | |
| Resolve voice string -> (preset_key, instruct). | |
| Ưu tiên: voice trực tiếp là preset -> dùng luôn. | |
| Nếu voice là extended (meme_male_gasp...) -> lấy base preset. | |
| Nếu style override -> dùng style để chọn tag. | |
| Returns (preset_key, instruct, is_extended) | |
| """ | |
| v = str(voice or "").strip() | |
| s = str(style or "").strip().lower() | |
| # direct preset | |
| if v in OMNIVOICE_PRESETS: | |
| return v, OMNIVOICE_PRESETS[v], False | |
| if v in OMNIVOICE_EXTENDED: | |
| base, _ = OMNIVOICE_EXTENDED[v] | |
| return base, OMNIVOICE_PRESETS.get(base, "male, young adult, moderate pitch"), True | |
| # fallback: try prefix match | |
| for preset in OMNIVOICE_PRESETS: | |
| if v.startswith(preset): | |
| return preset, OMNIVOICE_PRESETS[preset], True | |
| # default | |
| return "meme_male", OMNIVOICE_PRESETS["meme_male"], False | |
| def _omnivoice_tag_for_text(text, style="", voice=""): | |
| """ | |
| Determine OmniVoice emotion tag with priority: | |
| 1. Explicit [tag] prefix in text (highest) | |
| 2. Explicit style parameter (if not "auto"/"default"/empty) | |
| 3. Style inferred from voice name (e.g., "meme_male_gasp" -> "gasp") | |
| 4. Auto-detect from text content (lowest, only if no explicit style AND not pure preset) | |
| FIX 2026-08-29: Nếu user chọn 1 preset thuần (fun_male/story_male/meme_male/calm_fun_male) + style auto, | |
| khóa tag = "normal" để không nhảy giọng lung tung theo text. | |
| """ | |
| text = str(text or "").strip() | |
| style_l = str(style or "").strip().lower() | |
| voice_l = str(voice or "").strip().lower() | |
| tag = "" | |
| # 1. Explicit [tag] prefix in text (e.g., "[gasp] text") | |
| m = re.match(r"^\s*\[([^\]]+)\]\s*", text) | |
| if m: | |
| tag = m.group(1).strip().lower() | |
| # 2. Explicit style parameter (user-selected in GUI) - HIGHEST PROGRAMMATIC PRIORITY | |
| # Treat "auto" as "use voice default" not "skip to text detection" | |
| if not tag and style_l and style_l not in {"", "default", "auto"}: | |
| tag = style_l | |
| # 3. Infer from voice name if no explicit style (e.g., "meme_male_gasp" -> "gasp") | |
| if not tag: | |
| for candidate in ["gasp", "excited", "sad", "whispering", "sarcastic", "confident", "playful", "normal"]: | |
| if candidate in voice_l: | |
| tag = candidate | |
| break | |
| # Normalize tag aliases | |
| if tag in {"gasps", "gasping", "wow"}: | |
| tag = "gasp" | |
| # 4. Auto-detect from text content ONLY if no explicit style/voice style was set | |
| # FIX: Nếu voice là preset thuần (fun_male...) + style auto -> khóa normal, không auto theo text | |
| if not tag: | |
| # Kiểm tra pure preset -> lock | |
| if voice_l in OMNIVOICE_PRESETS and style_l in ("", "auto", "default"): | |
| tag = "normal" | |
| else: | |
| text_lower = text.lower() | |
| if any(w in text_lower for w in ["ơi", "a ha", "đùa", "vui", "haha", "quá đã", "cháy"]): | |
| tag = "excited" | |
| elif any(w in text_lower for w in ["buồn", "khóc", "đau lòng", "tiếc", "haizz"]): | |
| tag = "sad" | |
| elif any(w in text_lower for w in ["suỵt", "nói nhỏ", "thầm", "bí mật"]): | |
| tag = "whispering" | |
| elif any(w in text_lower for w in ["tin được không", "bất ngờ", "cái gì", "wow"]): | |
| tag = "gasp" | |
| elif any(w in text_lower for w in ["mỉa mai", "thế cơ à", "vậy hả", "chắc chưa"]): | |
| tag = "sarcastic" | |
| elif any(w in text_lower for w in ["chắc chắn", "tự tin", "khẳng định", "luôn"]): | |
| tag = "confident" | |
| else: | |
| tag = "normal" | |
| # Strip [tag] prefix from text for actual synthesis | |
| clean_text = re.sub(r"^\s*\[[^\]]+\]\s*", "", text).strip() | |
| return tag, clean_text or text | |
| def _omnivoice_params(tag, tts_config): | |
| # Base presets đồng bộ với omnivoice-web + emotion overrides | |
| # Khi voice là fun_male/story_male/... ta sẽ ưu tiên instruct của preset đó | |
| # trước rồi mới apply tag-based speed/guidance tweak | |
| preset = { | |
| "normal": ("male, young adult, moderate pitch", 1.06, 20, 2.0), | |
| "gasp": ("male, young adult, high pitch", 1.15, 20, 2.2), | |
| "excited": ("male, young adult, high pitch", 1.16, 20, 2.3), | |
| "sad": ("male, young adult, low pitch", 0.90, 20, 1.8), | |
| "whispering": ("male, young adult, whisper", 0.95, 20, 1.6), | |
| "sarcastic": ("male, young adult, moderate pitch", 0.98, 20, 2.0), | |
| "confident": ("male, young adult, low pitch", 1.02, 20, 2.0), | |
| "playful": ("male, young adult, high pitch", 1.08, 20, 2.1), | |
| } | |
| # Nếu tts_config có omnivoice_preset (fun_male...) ưu tiên instruct preset đó | |
| base_preset_key = tts_config.get("omnivoice_preset", "") | |
| base_instruct = OMNIVOICE_PRESETS.get(base_preset_key) | |
| instruct, speed, steps, guidance = preset.get(tag, preset["normal"]) | |
| # Override instruct nếu có base preset | |
| if base_instruct: | |
| instruct = base_instruct | |
| # tweak speed/guidance theo tag nhưng giữ instruct của preset | |
| if tag == "gasp": | |
| speed, guidance = 1.15, 2.2 | |
| elif tag == "excited": | |
| speed, guidance = 1.16, 2.3 | |
| elif tag == "sad": | |
| speed, guidance = 0.90, 1.8 | |
| elif tag == "whispering": | |
| speed, guidance = 0.95, 1.6 | |
| return { | |
| "instruct": tts_config.get("omnivoice_instruct", instruct), | |
| "speed": float(tts_config.get("omnivoice_speed", speed)), | |
| "steps": int(tts_config.get("omnivoice_steps", steps)), | |
| "guidance": float(tts_config.get("omnivoice_guidance", guidance)), | |
| } | |
| class OmniVoiceSession: | |
| def __init__(self, args, tts_config, work_dir): | |
| self.args = args | |
| self.tts_config = tts_config | |
| self.work_dir = Path(work_dir) | |
| self.process = None | |
| self.queue = queue.Queue() | |
| self.reader_thread = None | |
| self.runner_path = self.work_dir / "omnivoice_session_runner.py" | |
| def start(self): | |
| python_exe = Path(self.tts_config.get("omnivoice_python", r"C:\Users\Admin\OmniVoiceApp\.venv\Scripts\python.exe")) | |
| if not python_exe.exists(): | |
| raise RuntimeError(f"OmniVoice python not found: {python_exe}") | |
| model_name = self.tts_config.get("omnivoice_model", "splendor1811/omnivoice-vietnamese") | |
| ref_audio = self.tts_config.get("omnivoice_ref_audio", "") | |
| ref_text = self.tts_config.get("omnivoice_ref_text", "") | |
| runner_code = f"""import sys | |
| import json | |
| import traceback | |
| from pathlib import Path | |
| from omnivoice import OmniVoice | |
| import soundfile as sf | |
| import torch | |
| def emit(payload): | |
| print(json.dumps(payload, ensure_ascii=False), flush=True) | |
| try: | |
| device = "cuda:0" if torch.cuda.is_available() else "cpu" | |
| dtype = torch.float16 if torch.cuda.is_available() else torch.float32 | |
| model = OmniVoice.from_pretrained( | |
| {model_name!r}, | |
| device_map=device, | |
| dtype=dtype, | |
| ) | |
| ref_audio = {ref_audio!r} | |
| ref_text = {ref_text!r} | |
| voice_prompt = None | |
| if ref_audio and Path(ref_audio).exists(): | |
| if not ref_text: | |
| from transformers import pipeline as hf_pipeline | |
| asr_dtype = torch.float16 if torch.cuda.is_available() else torch.float32 | |
| asr_pipe = hf_pipeline( | |
| "automatic-speech-recognition", | |
| model="openai/whisper-large-v3-turbo", | |
| dtype=asr_dtype, | |
| device_map=device, | |
| ) | |
| ref_text = asr_pipe(ref_audio)["text"].strip() | |
| voice_prompt = model.create_voice_clone_prompt(ref_audio=ref_audio, ref_text=ref_text) | |
| emit({{"type": "clone_ready", "ok": True, "ref_text": ref_text}}) | |
| emit({{"type": "ready", "device": device}}) | |
| for line in sys.stdin: | |
| line = line.strip() | |
| if not line: | |
| continue | |
| try: | |
| req = json.loads(line) | |
| if req.get("type") == "quit": | |
| emit({{"type": "quit_ok"}}) | |
| break | |
| output = Path(req["output"]) | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| gen_kw = dict( | |
| text=req["text"], | |
| language="vi", | |
| speed=float(req["speed"]), | |
| num_step=int(req["steps"]), | |
| guidance_scale=float(req["guidance"]), | |
| denoise=True, | |
| ) | |
| if voice_prompt is not None: | |
| gen_kw["voice_clone_prompt"] = voice_prompt | |
| else: | |
| gen_kw["instruct"] = req["instruct"] | |
| audio = model.generate(**gen_kw) | |
| sf.write(str(output), audio[0], model.sampling_rate) | |
| emit({{"type": "result", "id": req.get("id"), "ok": True}}) | |
| except Exception as e: | |
| emit({{"type": "result", "id": req.get("id") if 'req' in locals() else None, "ok": False, "error": str(e), "traceback": traceback.format_exc()}}) | |
| except Exception as e: | |
| emit({{"type": "fatal", "ok": False, "error": str(e), "traceback": traceback.format_exc()}}) | |
| sys.exit(1) | |
| """ | |
| self.runner_path.write_text(runner_code, encoding="utf-8") | |
| startupinfo = None | |
| if sys.platform == 'win32': | |
| startupinfo = subprocess.STARTUPINFO() | |
| startupinfo.dwFlags |= subprocess.STARTF_USESHOWWINDOW | |
| self.process = subprocess.Popen( | |
| [str(python_exe), str(self.runner_path)], | |
| stdin=subprocess.PIPE, | |
| stdout=subprocess.PIPE, | |
| stderr=subprocess.STDOUT, | |
| text=True, | |
| encoding="utf-8", | |
| errors="ignore", | |
| startupinfo=startupinfo, | |
| bufsize=1, | |
| ) | |
| self.reader_thread = threading.Thread(target=self._reader_loop, daemon=True) | |
| self.reader_thread.start() | |
| ready_timeout = int(self.tts_config.get("omnivoice_session_ready_timeout_seconds", 300)) | |
| ready = self._wait_for(lambda payload: payload.get("type") == "ready", ready_timeout) | |
| if not ready: | |
| self.close(kill=True) | |
| raise RuntimeError("OmniVoice persistent session did not become ready before timeout.") | |
| print(f"[OMNIVOICE SESSION] ready device={ready.get('device')}") | |
| def _reader_loop(self): | |
| try: | |
| for line in iter(self.process.stdout.readline, ''): | |
| line = line.strip() | |
| if not line: | |
| continue | |
| try: | |
| payload = json.loads(line) | |
| self.queue.put(payload) | |
| except Exception: | |
| print(f"[OMNIVOICE SESSION] {line}") | |
| finally: | |
| self.queue.put({"type": "process_exit"}) | |
| def _wait_for(self, predicate, timeout_seconds): | |
| deadline = time.time() + max(1, timeout_seconds) | |
| while time.time() < deadline: | |
| remaining = max(0.1, deadline - time.time()) | |
| try: | |
| payload = self.queue.get(timeout=min(1.0, remaining)) | |
| except queue.Empty: | |
| if self.process and self.process.poll() is not None: | |
| return None | |
| continue | |
| if payload.get("type") == "fatal": | |
| raise RuntimeError(payload.get("error", "OmniVoice session fatal error")) | |
| if payload.get("type") == "process_exit": | |
| return None | |
| if predicate(payload): | |
| return payload | |
| print(f"[OMNIVOICE SESSION] {payload}") | |
| return None | |
| def synthesize(self, text, output_path, params): | |
| if not self.process or self.process.poll() is not None: | |
| raise RuntimeError("OmniVoice persistent session is not running.") | |
| req_id = uuid.uuid4().hex | |
| request = { | |
| "type": "synthesize", | |
| "id": req_id, | |
| "text": text, | |
| "output": str(Path(output_path).resolve()), | |
| "instruct": params["instruct"], | |
| "speed": params["speed"], | |
| "steps": params["steps"], | |
| "guidance": params["guidance"], | |
| } | |
| self.process.stdin.write(json.dumps(request, ensure_ascii=False) + "\n") | |
| self.process.stdin.flush() | |
| timeout = int(self.tts_config.get("omnivoice_timeout_seconds", 240)) | |
| result = self._wait_for(lambda payload: payload.get("type") == "result" and payload.get("id") == req_id, timeout) | |
| if not result: | |
| raise RuntimeError("OmniVoice persistent session timed out for segment.") | |
| if not result.get("ok"): | |
| raise RuntimeError(result.get("error", "OmniVoice persistent session failed")) | |
| def close(self, kill=False): | |
| try: | |
| if self.process and self.process.poll() is None: | |
| if kill: | |
| self.process.kill() | |
| else: | |
| try: | |
| self.process.stdin.write(json.dumps({"type": "quit"}) + "\n") | |
| self.process.stdin.flush() | |
| self.process.wait(timeout=10) | |
| except Exception: | |
| self.process.kill() | |
| finally: | |
| try: | |
| self.runner_path.unlink() | |
| except Exception: | |
| pass | |
| def _omnivoice_request_params(text, args, tts_config): | |
| """ | |
| KHÓA GIỌNG 100%: giọng nào khóa giọng đó, không biến đổi pitch theo text. | |
| Chỉ cho phép tăng tốc độ (speed) chung qua args.speed + ffmpeg atempo (ở phần time-stretch sau), | |
| không tăng/giảm tone. Tag chỉ lấy từ voice/style hoặc [tag] tường minh, không auto-detect theo text. | |
| """ | |
| preset_key, preset_instruct, is_extended = _resolve_omnivoice_preset(args.voice, args.style) | |
| cfg = dict(tts_config) | |
| cfg["omnivoice_preset"] = preset_key | |
| # Determine locked tag từ voice/style, bỏ auto-detect theo text | |
| voice_str = str(args.voice or "").strip().lower() | |
| style_str = str(args.style or "").strip().lower() | |
| text_str = str(text or "").strip() | |
| # Check explicit [tag] prefix — vẫn tôn trọng nếu user cố tình viết | |
| m = re.match(r"^\s*\[([^\]]+)\]\s*", text_str) | |
| explicit_tag = m.group(1).strip().lower() if m else "" | |
| if explicit_tag: | |
| # Normalize alias | |
| if explicit_tag in {"gasps", "gasping", "wow"}: | |
| explicit_tag = "gasp" | |
| tag = explicit_tag | |
| clean_text = re.sub(r"^\s*\[[^\]]+\]\s*", "", text_str).strip() or text_str | |
| elif voice_str in OMNIVOICE_PRESETS and style_str in ("", "auto", "default"): | |
| # Pure preset + auto => khóa normal (không nhảy) | |
| tag = "normal" | |
| clean_text = re.sub(r"^\s*\[[^\]]+\]\s*", "", text_str).strip() or text_str | |
| elif voice_str in OMNIVOICE_EXTENDED: | |
| # Extended đã khóa sẵn tag trong tên (meme_male_gasp -> gasp) | |
| tag = OMNIVOICE_EXTENDED.get(str(args.voice).strip(), (preset_key, "normal"))[1] | |
| clean_text = re.sub(r"^\s*\[[^\]]+\]\s*", "", text_str).strip() or text_str | |
| elif style_str and style_str not in ("", "auto", "default"): | |
| # Style tường minh (gasp/excited...) => khóa theo style | |
| tag = style_str | |
| if tag in {"gasps", "gasping", "wow"}: | |
| tag = "gasp" | |
| clean_text = re.sub(r"^\s*\[[^\]]+\]\s*", "", text_str).strip() or text_str | |
| else: | |
| # Fallback: voice không phải Omni preset (edge/piper) thì mới cho auto-detect | |
| # Nhưng với Omni thì vẫn khóa normal để tránh nhảy | |
| if voice_str in OMNIVOICE_PRESETS or any(voice_str.startswith(p) for p in OMNIVOICE_PRESETS): | |
| tag = "normal" | |
| clean_text = re.sub(r"^\s*\[[^\]]+\]\s*", "", text_str).strip() or text_str | |
| else: | |
| tag, clean_text = _omnivoice_tag_for_text(text_str, style=args.style, voice=args.voice) | |
| base_tag_params = _omnivoice_params(tag, cfg) | |
| # Ép instruct khóa theo preset đã chọn (không cho tag đổi tone) | |
| # Với pure preset -> preset_instruct, với extended cũng vậy (đã là base preset) | |
| base_tag_params["instruct"] = preset_instruct | |
| # Speed khóa theo tag của voice (đã cố định ở trên) * global multiplier — không auto theo text | |
| params = base_tag_params | |
| try: | |
| speed_multiplier = float(args.speed) | |
| if speed_multiplier > 0: | |
| # Chỉ nhân global speed, không để tag làm nhảy thêm | |
| # base speed đã là của locked tag (normal 1.06, gasp 1.15...), nhân thêm global | |
| params["speed"] = max(0.75, min(1.65, params["speed"] * speed_multiplier / _omnivoice_params(tag, {"omnivoice_preset": preset_key})["speed"] * params["speed"])) # keep as is, simplified below | |
| # Actually simpler: params speed already includes tag speed, just multiply global | |
| # Đã làm ở trên, nhưng để tránh double, ta tính lại: lấy base speed của locked tag nhân global | |
| # Do base_tag_params đã có tag speed, ta chỉ cần nhân global | |
| pass | |
| except Exception: | |
| pass | |
| # Thực hiện nhân global đúng: base_tag_params đã có speed của locked tag, nhân thêm | |
| try: | |
| gm = float(args.speed) | |
| if gm and gm != 1.0: | |
| # base_tag_params speed hiện là speed của locked tag, cần nhân gm (đã làm ở trên nhưng giữ cho chắc) | |
| # Nếu gm !=1, params speed đã nhân ở trên, bỏ qua | |
| pass | |
| except: | |
| pass | |
| # Chuẩn hóa lần cuối: speed = locked_tag_speed * global | |
| # Để tránh nhân đôi, lấy lại locked speed gốc | |
| try: | |
| locked_base = _omnivoice_params(tag, {"omnivoice_preset": preset_key}) | |
| gm = float(args.speed) if str(args.speed).replace('.','',1).replace('-','',1).isdigit() else 1.0 | |
| params["speed"] = max(0.75, min(1.65, locked_base["speed"] * gm)) | |
| params["instruct"] = preset_instruct # khóa cứng | |
| except: | |
| pass | |
| return tag, clean_text, params | |
| def run_omnivoice_tts(text, output_path, args, tts_config, omnivoice_session=None): | |
| python_exe = Path(tts_config.get("omnivoice_python", r"C:\Users\Admin\OmniVoiceApp\.venv\Scripts\python.exe")) | |
| if not python_exe.exists(): | |
| raise RuntimeError(f"OmniVoice python not found: {python_exe}") | |
| tag, clean_text, params = _omnivoice_request_params(text, args, tts_config) | |
| output_path = Path(output_path) | |
| output_path.parent.mkdir(parents=True, exist_ok=True) | |
| print( | |
| f"[OMNIVOICE] tag={tag} instruct={params['instruct']} " | |
| f"speed={params['speed']:.2f} steps={params['steps']} guidance={params['guidance']}" | |
| ) | |
| if omnivoice_session: | |
| omnivoice_session.synthesize(clean_text, output_path, params) | |
| return | |
| runner_path = output_path.parent / f"omnivoice_runner_{output_path.stem}.py" | |
| model_name = tts_config.get("omnivoice_model", "splendor1811/omnivoice-vietnamese") | |
| ref_audio = tts_config.get("omnivoice_ref_audio", "") | |
| ref_text = tts_config.get("omnivoice_ref_text", "") | |
| runner_code = f"""import sys | |
| from pathlib import Path | |
| from omnivoice import OmniVoice | |
| import soundfile as sf | |
| import torch | |
| try: | |
| device = "cuda:0" if torch.cuda.is_available() else "cpu" | |
| dtype = torch.float16 if torch.cuda.is_available() else torch.float32 | |
| model = OmniVoice.from_pretrained( | |
| {model_name!r}, | |
| device_map=device, | |
| dtype=dtype, | |
| ) | |
| ref_audio = {ref_audio!r} | |
| ref_text = {ref_text!r} | |
| voice_prompt = None | |
| if ref_audio and Path(ref_audio).exists(): | |
| if not ref_text: | |
| from transformers import pipeline as hf_pipeline | |
| asr_dtype = torch.float16 if torch.cuda.is_available() else torch.float32 | |
| asr_pipe = hf_pipeline( | |
| "automatic-speech-recognition", | |
| model="openai/whisper-large-v3-turbo", | |
| dtype=asr_dtype, | |
| device_map=device, | |
| ) | |
| ref_text = asr_pipe(ref_audio)["text"].strip() | |
| voice_prompt = model.create_voice_clone_prompt(ref_audio=ref_audio, ref_text=ref_text) | |
| gen_kw = dict( | |
| text={clean_text!r}, | |
| language="vi", | |
| speed={params["speed"]}, | |
| num_step={params["steps"]}, | |
| guidance_scale={params["guidance"]}, | |
| denoise=True, | |
| ) | |
| if voice_prompt is not None: | |
| gen_kw["voice_clone_prompt"] = voice_prompt | |
| else: | |
| gen_kw["instruct"] = {params["instruct"]!r} | |
| audio = model.generate(**gen_kw) | |
| sf.write({str(output_path.resolve())!r}, audio[0], model.sampling_rate) | |
| print("SUCCESS") | |
| except Exception as e: | |
| print("ERROR:", e) | |
| sys.exit(1) | |
| """ | |
| runner_path.write_text(runner_code, encoding="utf-8") | |
| startupinfo = None | |
| if sys.platform == 'win32': | |
| startupinfo = subprocess.STARTUPINFO() | |
| startupinfo.dwFlags |= subprocess.STARTF_USESHOWWINDOW | |
| try: | |
| result = subprocess.run( | |
| [str(python_exe), str(runner_path)], | |
| capture_output=True, | |
| text=True, | |
| encoding="utf-8", | |
| errors="ignore", | |
| timeout=int(tts_config.get("omnivoice_timeout_seconds", 240)), | |
| startupinfo=startupinfo, | |
| ) | |
| if result.returncode != 0: | |
| raise RuntimeError((result.stderr or result.stdout or "OmniVoice failed").strip()) | |
| finally: | |
| try: | |
| runner_path.unlink() | |
| except Exception: | |
| pass | |
| def run_omnivoice_cloud_tts(text, output_path, args, tts_config): | |
| """ | |
| Goi OmniVoice Cloud qua Gradio (primary, đã verify OK) + FastAPI fallback. | |
| Dong bo 100% voi omnivoice-web PRESETS. Gradio api_name="/gradio_fn" đã test trả audio.wav. | |
| """ | |
| api_url = tts_config.get("omnivoice_api_url", "") or tts_config.get("omnivoice_cloud_url", "") or os.getenv("OMNIVOICE_API_URL", "https://hoangtaiii-omnivoice.hf.space") | |
| api_key = tts_config.get("omnivoice_api_key", "") or tts_config.get("omnivoice_cloud_api_key", "") or os.getenv("OMNIVOICE_API_KEY", "sk-demo123") | |
| api_url = str(api_url).rstrip("/") | |
| preset_key, _, _ = _resolve_omnivoice_preset(args.voice, args.style) | |
| preset = preset_key if preset_key in OMNIVOICE_PRESETS else "meme_male" | |
| try: | |
| speed_val = float(args.speed) if str(args.speed).replace('.','',1).isdigit() or str(args.speed).replace('.','',1).replace('-','',1).isdigit() else 1.08 | |
| except: | |
| speed_val = 1.08 | |
| speed_val = max(0.5, min(1.5, speed_val)) | |
| output_path = Path(output_path) | |
| output_path.parent.mkdir(parents=True, exist_ok=True) | |
| # 1) Thử Gradio trước — truyền HF_TOKEN để tăng ZeroGPU quota | |
| gradio_err = None | |
| try: | |
| from gradio_client import Client | |
| hf_tok = tts_config.get("hf_token","") or tts_config.get("hf_key","") or tts_config.get("HUGGING_FACE_HUB_TOKEN","") or os.getenv("HF_TOKEN","") or os.getenv("HUGGING_FACE_HUB_TOKEN","") or os.getenv("HUGGINGFACE_TOKEN","") | |
| # Fallback đọc hf_key từ config.json root | |
| if not hf_tok: | |
| try: | |
| _cfg_path = Path(__file__).resolve().parents[2] / "config.json" | |
| if _cfg_path.exists(): | |
| _j = json.loads(_cfg_path.read_text(encoding="utf-8")) | |
| hf_tok = _j.get("hf_key","") or _j.get("hf_token","") or hf_tok | |
| except: pass | |
| if hf_tok: | |
| client = Client(api_url, token=hf_tok.strip()) | |
| print(f"[OMNIVOICE CLOUD] Gradio + HF_TOKEN {hf_tok[:6]}... {api_url} preset={preset} speed={speed_val} len={len(str(text))}") | |
| else: | |
| client = Client(api_url) | |
| print(f"[OMNIVOICE CLOUD] Gradio (anonymous, quota thấp) {api_url} preset={preset} speed={speed_val} len={len(str(text))} — set hf_key/HF_TOKEN để tăng quota") | |
| result = client.predict(str(text).strip(), preset, float(speed_val), 16, 2.0, api_name="/gradio_fn") | |
| src_path = None | |
| if isinstance(result, str) and Path(result).exists(): | |
| src_path = Path(result) | |
| elif isinstance(result, dict) and "path" in result and Path(result["path"]).exists(): | |
| src_path = Path(result["path"]) | |
| elif isinstance(result, (list, tuple)) and len(result) > 0: | |
| cand = result[0] | |
| if isinstance(cand, str) and Path(cand).exists(): | |
| src_path = Path(cand) | |
| elif isinstance(cand, dict) and "path" in cand and Path(cand["path"]).exists(): | |
| src_path = Path(cand["path"]) | |
| # Gradio 6 có thể trả tuple (sr, wav_array) hoặc filepath — handle thêm | |
| if src_path and src_path.exists(): | |
| # Dùng ffmpeg chuẩn hóa 16k mono (giữ nguyên nếu đã 16k) | |
| # Nếu src là wav 24k của OmniVoice -> convert | |
| import subprocess | |
| # Lấy ffmpeg path từ args nếu có | |
| ff = getattr(args, "ffmpeg_path", "ffmpeg") | |
| if not ff: | |
| ff = tts_config.get("ffmpeg_path", "ffmpeg") or "ffmpeg" | |
| # Dùng pydub verify sau convert | |
| cmd = [str(ff), "-y", "-i", str(src_path), "-ar", "16000", "-ac", "1", str(output_path)] | |
| # Dùng startupinfo ẩn cửa sổ trên Windows | |
| startupinfo = None | |
| if sys.platform == 'win32': | |
| startupinfo = subprocess.STARTUPINFO() | |
| startupinfo.dwFlags |= subprocess.STARTF_USESHOWWINDOW | |
| subprocess.run(cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=True, startupinfo=startupinfo) | |
| if not verify_audio_file(output_path): | |
| raise RuntimeError("Gradio audio verify failed (<500 bytes)") | |
| print(f"[OMNIVOICE CLOUD] Gradio OK -> {output_path} ({output_path.stat().st_size} bytes)") | |
| return | |
| else: | |
| raise RuntimeError(f"Gradio result not file: {result!r:.800}") | |
| except Exception as e: | |
| gradio_err = e | |
| print(f"[OMNIVOICE CLOUD] Gradio failed: {e} -> fallback FastAPI") | |
| # 2) Fallback FastAPI | |
| import requests | |
| payload = { | |
| "text": str(text).strip(), | |
| "preset": preset, | |
| "speed": float(speed_val), | |
| "steps": 16, | |
| "guidance": 2.0, | |
| "language": "vi" | |
| } | |
| if tts_config.get("omnivoice_instruct"): | |
| payload["instruct"] = tts_config["omnivoice_instruct"] | |
| headers = {"X-API-Key": api_key, "Content-Type": "application/json"} | |
| print(f"[OMNIVOICE CLOUD] Fallback POST {api_url}/v1/generate preset={preset} len={len(payload['text'])}") | |
| resp = requests.post(f"{api_url}/v1/generate", headers=headers, json=payload, timeout=int(tts_config.get("omnivoice_timeout_seconds", 240))) | |
| if resp.status_code == 404 and gradio_err is not None: | |
| raise RuntimeError(f"FastAPI 404 + Gradio failed ({gradio_err}): {resp.text[:300]}") | |
| if resp.status_code != 200: | |
| raise RuntimeError(f"OmniVoice Cloud {resp.status_code}: {resp.text[:500]} (Gradio err: {gradio_err})") | |
| ctype = resp.headers.get("content-type", "") | |
| if "audio" not in ctype and len(resp.content) < 1000: | |
| try: | |
| j = resp.json() | |
| raise RuntimeError(f"OmniVoice Cloud JSON error: {j} (Gradio err: {gradio_err})") | |
| except Exception: | |
| pass | |
| output_path.write_bytes(resp.content) | |
| if not verify_audio_file(output_path): | |
| raise RuntimeError(f"OmniVoice Cloud invalid audio (<500) Gradio err: {gradio_err}") | |
| print(f"[OMNIVOICE CLOUD] FastAPI fallback OK -> {output_path} ({output_path.stat().st_size} bytes)") | |
| def synthesize_text_to_wav(text, output_path, args, tts_config, piper_voice=None, voice_override=None, omnivoice_session=None): | |
| engine = str(args.engine or "").lower().strip() | |
| # Chuẩn hóa alias: "omnivoice (local)" -> "omnivoice", "omnivoice cloud (hf space)" -> "omnivoice_cloud" | |
| if "cloud" in engine or "hf" in engine or "api" in engine: | |
| engine = "omnivoice_cloud" | |
| elif "omnivoice" in engine: | |
| engine = "omnivoice" | |
| voice = voice_override or args.voice | |
| if engine == "piper" and piper_voice: | |
| import wave | |
| from piper import SynthesisConfig | |
| with wave.open(str(output_path), "wb") as wav_file: | |
| try: | |
| speed_val = float(args.speed) | |
| length_scale = 1.0 / speed_val | |
| except Exception: | |
| length_scale = 1.0 | |
| syn_config = SynthesisConfig(length_scale=length_scale) | |
| piper_voice.synthesize_wav(text, wav_file, syn_config=syn_config) | |
| elif engine == "omnivoice": | |
| run_omnivoice_tts(text, output_path, args, tts_config, omnivoice_session=omnivoice_session) | |
| elif engine == "omnivoice_cloud": | |
| # Nếu args.voice bị override bởi fallback, tạm gán lại args.voice = voice để resolve preset đúng | |
| orig_voice = args.voice | |
| try: | |
| if voice_override: | |
| args.voice = voice_override | |
| run_omnivoice_cloud_tts(text, output_path, args, tts_config) | |
| finally: | |
| args.voice = orig_voice | |
| else: | |
| asyncio.run(run_edge_tts(text, voice, output_path, args.speed, args.pitch, args.volume)) | |
| def _block_should_generate_tts(block, action_by_id, labels, skipped_tts_ids): | |
| b_id = str(block.get("id")) | |
| action = action_by_id.get(b_id, {}) | |
| should_generate_tts = action.get("generate_vi_voiceover", True) | |
| if labels.get(b_id) == "EN_PRESERVE" or b_id in skipped_tts_ids or not should_generate_tts: | |
| return False | |
| detected_language = action.get("detected_language", labels.get(b_id, "unknown")) | |
| if detected_language in {"en", "music", "unknown", "mixed"} and not should_generate_tts: | |
| return False | |
| return True | |
| def _normalize_tts_text(text): | |
| return re.sub(r"\s+", " ", re.sub(r"[^\wÀ-ỹ]+", " ", str(text or "").lower(), flags=re.UNICODE)).strip() | |
| def _join_tts_text(left, right): | |
| left = str(left or "").strip() | |
| right = str(right or "").strip() | |
| if not left: | |
| return right | |
| if not right: | |
| return left | |
| left_key = _normalize_tts_text(left) | |
| right_key = _normalize_tts_text(right) | |
| if left_key == right_key: | |
| return left | |
| if left_key and right_key.startswith(left_key): | |
| return right | |
| if right_key and left_key.endswith(right_key): | |
| return left | |
| parts = [] | |
| seen = set() | |
| for item in re.split(r"(?<=[.!?])\s+", f"{left} {right}"): | |
| item = item.strip(" ,") | |
| key = _normalize_tts_text(item) | |
| if not key or key in seen: | |
| continue | |
| parts.append(item) | |
| seen.add(key) | |
| return " ".join(parts).strip() | |
| def _merge_short_tts_blocks(blocks, action_by_id, labels, skipped_tts_ids, tts_config): | |
| enabled = bool(tts_config.get("merge_short_segments_enabled", False)) | |
| min_input_blocks = int(tts_config.get("merge_short_segments_min_input_blocks", 80)) | |
| if not enabled or len(blocks) < min_input_blocks: | |
| return blocks, { | |
| "enabled": False, | |
| "reason": "disabled_or_too_few_blocks", | |
| "input_blocks": len(blocks), | |
| "output_units": len(blocks), | |
| } | |
| max_gap_ms = int(tts_config.get("merge_short_segments_max_gap_ms", 260)) | |
| short_gap_ms = int(tts_config.get("merge_short_segments_short_gap_ms", 900)) | |
| target_min_duration_ms = int(tts_config.get("merge_short_segments_target_min_duration_ms", 1600)) | |
| max_duration_ms = int(tts_config.get("merge_short_segments_max_duration_ms", 4200)) | |
| max_chars = int(tts_config.get("merge_short_segments_max_chars", 86)) | |
| max_blocks = int(tts_config.get("merge_short_segments_max_blocks", 5)) | |
| units = [] | |
| current = None | |
| def flush(): | |
| nonlocal current | |
| if not current: | |
| return | |
| ids = current["source_ids"] | |
| current["id"] = ids[0] | |
| current["merged"] = len(ids) > 1 | |
| current["merged_block_count"] = len(ids) | |
| units.append(current) | |
| current = None | |
| for block in blocks: | |
| b = dict(block) | |
| b["source_ids"] = [str(block.get("id"))] | |
| b["merged"] = False | |
| b["merged_block_count"] = 1 | |
| if not _block_should_generate_tts(b, action_by_id, labels, skipped_tts_ids): | |
| flush() | |
| units.append(b) | |
| continue | |
| text = str(b.get("text", "")).strip() | |
| if not current: | |
| current = b | |
| continue | |
| gap_ms = int(b["start_ms"]) - int(current["end_ms"]) | |
| combined_text = _join_tts_text(current.get("text", ""), text) | |
| combined_duration = int(b["end_ms"]) - int(current["start_ms"]) | |
| current_duration = int(current["end_ms"]) - int(current["start_ms"]) | |
| next_duration = int(b["end_ms"]) - int(b["start_ms"]) | |
| allowed_gap_ms = short_gap_ms if (current_duration < target_min_duration_ms or next_duration < target_min_duration_ms) else max_gap_ms | |
| can_merge = ( | |
| gap_ms <= allowed_gap_ms | |
| and combined_duration <= max_duration_ms | |
| and len(combined_text) <= max_chars | |
| and len(current["source_ids"]) < max_blocks | |
| ) | |
| if can_merge: | |
| current["end_ms"] = b["end_ms"] | |
| current["text"] = combined_text | |
| current["source_ids"].append(str(b.get("id"))) | |
| continue | |
| flush() | |
| current = b | |
| flush() | |
| merged_units = sum(1 for u in units if u.get("merged")) | |
| return units, { | |
| "enabled": True, | |
| "input_blocks": len(blocks), | |
| "output_units": len(units), | |
| "merged_units": merged_units, | |
| "saved_tts_calls": max(0, len(blocks) - len(units)), | |
| "max_gap_ms": max_gap_ms, | |
| "short_gap_ms": short_gap_ms, | |
| "target_min_duration_ms": target_min_duration_ms, | |
| "max_duration_ms": max_duration_ms, | |
| "max_chars": max_chars, | |
| "max_blocks": max_blocks, | |
| } | |
| def _tts_plan_signature(units): | |
| payload = [ | |
| { | |
| "id": str(u.get("id")), | |
| "source_ids": [str(x) for x in u.get("source_ids", [u.get("id")])], | |
| "start_ms": int(u.get("start_ms", 0)), | |
| "end_ms": int(u.get("end_ms", 0)), | |
| "text": str(u.get("text", "")), | |
| } | |
| for u in units | |
| ] | |
| raw = json.dumps(payload, ensure_ascii=False, sort_keys=True) | |
| return hashlib.sha256(raw.encode("utf-8")).hexdigest() | |
| def _prepare_tts_plan(segments_dir, units, merge_report): | |
| plan_path = Path(segments_dir).parent / "tts_merge_plan.json" | |
| signature = _tts_plan_signature(units) | |
| previous_signature = None | |
| if plan_path.exists(): | |
| try: | |
| previous_signature = json.loads(plan_path.read_text(encoding="utf-8")).get("signature") | |
| except Exception: | |
| previous_signature = None | |
| if previous_signature != signature: | |
| for wav_path in Path(segments_dir).glob("*.wav"): | |
| try: | |
| wav_path.unlink() | |
| except Exception: | |
| pass | |
| payload = { | |
| "signature": signature, | |
| "created_at": time.time(), | |
| "merge_report": merge_report, | |
| "units": [ | |
| { | |
| "id": str(u.get("id")), | |
| "source_ids": [str(x) for x in u.get("source_ids", [u.get("id")])], | |
| "start_ms": int(u.get("start_ms", 0)), | |
| "end_ms": int(u.get("end_ms", 0)), | |
| "merged": bool(u.get("merged", False)), | |
| "merged_block_count": int(u.get("merged_block_count", 1)), | |
| "text": str(u.get("text", "")), | |
| } | |
| for u in units | |
| ], | |
| } | |
| plan_path.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8") | |
| return plan_path, previous_signature != signature | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Standalone TTS Segment Synthesis Worker CLI") | |
| parser.add_argument("--srt", default="", help="Path to input translated SRT file") | |
| parser.add_argument("--original-srt", default="", help="Path to original (ASR/OCR) SRT file") | |
| parser.add_argument("--original-audio", default="", help="Path to original audio WAV file") | |
| parser.add_argument("--background-audio", default="", help="Path to background separated WAV file") | |
| parser.add_argument("--output-wav", default="", help="Path to output final dubbing WAV") | |
| parser.add_argument("--segments-dir", default="", help="Directory to save segment WAVs") | |
| parser.add_argument("--engine", default="edge-tts", help="TTS Engine: edge-tts or piper") | |
| parser.add_argument("--voice", default="vi-VN-HoaiMyNeural", help="Voice name to use") | |
| parser.add_argument("--ffmpeg-path", default="ffmpeg", help="Path to ffmpeg executable") | |
| parser.add_argument("--vocals-audio", default="", help="Path to vocals separated WAV file") | |
| parser.add_argument("--speed", default="1.0", help="Voice speed multiplier") | |
| parser.add_argument("--pitch", default="0", help="Voice pitch change (Hz or %)") | |
| parser.add_argument("--volume", default="100", help="Voice volume percentage (100 is default)") | |
| parser.add_argument("--style", default="", help="Voice style parameter") | |
| parser.add_argument("--language-segments", default="", help="Path to audio_language_segments.json") | |
| parser.add_argument("--test-mode", action="store_true", help="If true, silent fallback on errors") | |
| parser.add_argument("--debug-segment-text", default=None, help="Segment text to synthesize for debugging") | |
| parser.add_argument("--output", default=None, help="Output wav path for debugging") | |
| parser.add_argument("--limit-segments", type=int, default=0, help=argparse.SUPPRESS) | |
| args = parser.parse_args() | |
| # Load configuration | |
| config_path = Path(__file__).parent.parent.parent / "config.json" | |
| tts_config = {} | |
| if config_path.exists(): | |
| try: | |
| with open(config_path, "r", encoding="utf-8") as f: | |
| config_data = json.load(f) | |
| tts_config = config_data.get("tts", {}) | |
| except Exception as e: | |
| print(f"Warning: Failed to load config.json: {e}", file=sys.stderr) | |
| # Set AudioSegment converter early so verify_audio_file and silence placeholder work correctly | |
| from pydub import AudioSegment | |
| if args.ffmpeg_path: | |
| AudioSegment.converter = str(args.ffmpeg_path) | |
| if args.debug_segment_text: | |
| text = args.debug_segment_text | |
| output_path = Path(args.output) if args.output else Path("debug_tts_out.wav") | |
| print(f"[DEBUG TTS] Starting single segment debug for: '{text}'") | |
| print(f"[DEBUG TTS] Output path: {output_path}") | |
| print(f"[DEBUG TTS] Voice: {args.voice}") | |
| # We will run the plans one by one and print results | |
| # Plan 1: Normal | |
| print("Plan 1 (normal)...") | |
| success = False | |
| try: | |
| asyncio.run(run_edge_tts(text, args.voice, output_path, args.speed, args.pitch, args.volume)) | |
| if verify_audio_file(output_path): | |
| print("Plan 1: PASS") | |
| success = True | |
| else: | |
| print("Plan 1: FAIL (empty or invalid audio)") | |
| except Exception as e: | |
| print(f"Plan 1: FAIL ({e})") | |
| # Plan 2: Sanitized | |
| if not success: | |
| print("Plan 2 (sanitized)...") | |
| sanitized = sanitize_for_edge_tts(text) | |
| print(f"Sanitized text: '{sanitized}'") | |
| if not sanitized: | |
| print("Plan 2: SKIP (text is empty after sanitization)") | |
| else: | |
| try: | |
| asyncio.run(run_edge_tts(sanitized, args.voice, output_path, args.speed, args.pitch, args.volume)) | |
| if verify_audio_file(output_path): | |
| print("Plan 2: PASS") | |
| success = True | |
| else: | |
| print("Plan 2: FAIL (empty or invalid audio)") | |
| except Exception as e: | |
| print(f"Plan 2: FAIL ({e})") | |
| # Plan 3: Split | |
| if not success: | |
| print("Plan 3 (split)...") | |
| parts = split_text_in_half(text) | |
| print(f"Split parts: {parts}") | |
| if len(parts) >= 2: | |
| part1_wav = output_path.parent / "debug_part1.wav" | |
| part2_wav = output_path.parent / "debug_part2.wav" | |
| from pydub import AudioSegment | |
| AudioSegment.converter = str(args.ffmpeg_path) | |
| try: | |
| asyncio.run(run_edge_tts(parts[0], args.voice, part1_wav, args.speed, args.pitch, args.volume)) | |
| asyncio.run(run_edge_tts(parts[1], args.voice, part2_wav, args.speed, args.pitch, args.volume)) | |
| if verify_audio_file(part1_wav) and verify_audio_file(part2_wav): | |
| seg1 = AudioSegment.from_file(part1_wav) | |
| seg2 = AudioSegment.from_file(part2_wav) | |
| combined = seg1 + AudioSegment.silent(duration=100) + seg2 | |
| combined.export(str(output_path), format="wav") | |
| if verify_audio_file(output_path): | |
| print("Plan 3: PASS") | |
| success = True | |
| else: | |
| print("Plan 3: FAIL (empty combined)") | |
| else: | |
| print("Plan 3: FAIL (part 1 or part 2 failed)") | |
| except Exception as e: | |
| print(f"Plan 3: FAIL ({e})") | |
| finally: | |
| for p in [part1_wav, part2_wav]: | |
| if p.exists(): | |
| p.unlink() | |
| else: | |
| print("Plan 3: SKIP (cannot split)") | |
| # Plan 4: Retry cùng giọng user chọn (KHÔNG nhảy giọng khác) | |
| # [FIX] Khóa cứng voice — không dùng fallback voice để tránh nhiều giọng trong 1 video | |
| if not success: | |
| print(f"Plan 4 (retry same voice: {args.voice})...") | |
| try: | |
| asyncio.run(run_edge_tts(text, args.voice, output_path, args.speed, args.pitch, args.volume)) | |
| if verify_audio_file(output_path): | |
| print("Plan 4: PASS (same voice retry)") | |
| success = True | |
| else: | |
| print("Plan 4: FAIL (empty audio, will use silence)") | |
| except Exception as e: | |
| print(f"Plan 4: FAIL ({e}, will use silence)") | |
| # Plan 5: Silence placeholder | |
| if not success: | |
| print("Plan 5 (silence placeholder)...") | |
| try: | |
| from pydub import AudioSegment | |
| AudioSegment.converter = str(args.ffmpeg_path) | |
| AudioSegment.silent(duration=2000).export(str(output_path), format="wav") | |
| print("Plan 5: PASS") | |
| success = True | |
| except Exception as e: | |
| print(f"Plan 5: FAIL ({e})") | |
| print(f"[DEBUG TTS] Done. Success: {success}") | |
| sys.exit(0) | |
| # Validate non-debug args | |
| if not args.srt or not args.output_wav or not args.segments_dir: | |
| print("Error: --srt, --output-wav, and --segments-dir are required when not in debug mode.", file=sys.stderr) | |
| sys.exit(1) | |
| srt_path = Path(args.srt) | |
| original_srt_path = Path(args.original_srt) if args.original_srt else None | |
| original_audio_path = Path(args.original_audio) if args.original_audio else None | |
| background_audio_path = Path(args.background_audio) if args.background_audio else None | |
| vocals_audio_path = Path(args.vocals_audio) if args.vocals_audio else None | |
| output_wav = Path(args.output_wav) | |
| segments_dir = Path(args.segments_dir) | |
| ffmpeg_path = Path(args.ffmpeg_path) | |
| from pydub import AudioSegment | |
| AudioSegment.converter = str(ffmpeg_path) | |
| if not srt_path.exists(): | |
| print(f"Error: SRT file not found at {srt_path}", file=sys.stderr) | |
| sys.exit(1) | |
| segments_dir.mkdir(parents=True, exist_ok=True) | |
| normalizer = VietnameseTextNormalizer() | |
| # Load original blocks | |
| orig_blocks = [] | |
| if original_srt_path and original_srt_path.exists(): | |
| print(f"Parsing original SRT: {original_srt_path}") | |
| orig_blocks = parse_srt(original_srt_path) | |
| # Load language segments for classification if available | |
| language_segments = None | |
| if args.language_segments and Path(args.language_segments).exists(): | |
| try: | |
| with open(args.language_segments, "r", encoding="utf-8") as f: | |
| language_segments = json.load(f) | |
| print(f"Loaded language segments: {len(language_segments)} entries.") | |
| except Exception as e: | |
| print(f"Warning: Failed to load language segments: {e}", file=sys.stderr) | |
| # Classify blocks or load pre-computed regions | |
| preserve_intervals = [] | |
| skipped_tts_ids = [] | |
| labels = {} | |
| segment_actions = [] | |
| action_by_id = {} | |
| preserve_json = output_wav.parent / "preserve_regions.json" | |
| skipped_json = output_wav.parent / "tts_skipped_blocks.json" | |
| actions_json = output_wav.parent / "segment_actions.json" | |
| if preserve_json.exists() and skipped_json.exists(): | |
| try: | |
| with open(preserve_json, "r", encoding="utf-8") as f: | |
| preserve_intervals = [tuple(x) for x in json.load(f)] | |
| with open(skipped_json, "r", encoding="utf-8") as f: | |
| skipped_tts_ids = json.load(f) | |
| print(f"Loaded preserve intervals from cache: {len(preserve_intervals)} entries.") | |
| except Exception as e: | |
| print(f"Warning: Failed to load pre-computed regions: {e}", file=sys.stderr) | |
| if actions_json.exists(): | |
| try: | |
| with open(actions_json, "r", encoding="utf-8") as f: | |
| segment_actions = json.load(f) | |
| action_by_id = {str(x.get("id")): x for x in segment_actions} | |
| print(f"Loaded segment actions: {len(segment_actions)} entries.") | |
| except Exception as e: | |
| print(f"Warning: Failed to load segment actions: {e}", file=sys.stderr) | |
| if not preserve_intervals and not skipped_tts_ids: | |
| # Fallback to computing them | |
| classifier = AudioTimelineClassifier() | |
| labels, preserve_intervals, skipped_tts_ids, segment_actions = classifier.classify_blocks( | |
| orig_blocks, language_segments, include_actions=True | |
| ) | |
| action_by_id = {str(x.get("id")): x for x in segment_actions} | |
| try: | |
| with open(preserve_json, "w", encoding="utf-8") as f: | |
| json.dump(preserve_intervals, f, ensure_ascii=False, indent=2) | |
| with open(skipped_json, "w", encoding="utf-8") as f: | |
| json.dump(skipped_tts_ids, f, ensure_ascii=False, indent=2) | |
| with open(actions_json, "w", encoding="utf-8") as f: | |
| json.dump(segment_actions, f, ensure_ascii=False, indent=2) | |
| except Exception as e: | |
| print(f"Warning: Failed to save regions: {e}", file=sys.stderr) | |
| print(f"Parsing translated SRT: {srt_path}...") | |
| blocks = parse_srt(srt_path) | |
| if not blocks: | |
| print("Error: No valid SRT blocks found.", file=sys.stderr) | |
| sys.exit(1) | |
| input_block_count = len(blocks) | |
| blocks, merge_report = _merge_short_tts_blocks(blocks, action_by_id, labels, skipped_tts_ids, tts_config) | |
| plan_path, plan_changed = _prepare_tts_plan(segments_dir, blocks, merge_report) | |
| if merge_report.get("enabled"): | |
| print( | |
| "[TTS MERGE] " | |
| f"input_blocks={merge_report.get('input_blocks')} " | |
| f"output_units={merge_report.get('output_units')} " | |
| f"saved_tts_calls={merge_report.get('saved_tts_calls')}" | |
| ) | |
| else: | |
| print(f"[TTS MERGE] disabled: {merge_report.get('reason')}") | |
| if plan_changed: | |
| print(f"[TTS MERGE] plan changed; cleared stale segment WAVs. Plan: {plan_path}") | |
| if args.limit_segments and args.limit_segments > 0: | |
| blocks = blocks[:args.limit_segments] | |
| print(f"[DEBUG TTS] Limiting synthesis to first {len(blocks)} unit(s).") | |
| piper_voice = None | |
| if args.engine.lower() == "piper": | |
| print("Initializing Piper local TTS...") | |
| from piper import PiperVoice, SynthesisConfig | |
| voice_model_name = args.voice | |
| if not voice_model_name.startswith("vi_VN"): | |
| voice_model_name = "vi_VN-vais1000-medium" | |
| model_dir = Path(__file__).parent.parent.parent / "models" / "piper" | |
| try: | |
| onnx_path, json_path = download_piper_model(voice_model_name, model_dir) | |
| piper_voice = PiperVoice.load(str(onnx_path)) | |
| except Exception as e: | |
| print(f"Error loading Piper: {e}", file=sys.stderr) | |
| if not args.test_mode: | |
| sys.exit(3) | |
| omnivoice_session = None | |
| # Chỉ bật persistent session cho local OmniVoice; Cloud dùng HTTP không cần | |
| _engine_norm = str(args.engine or "").lower() | |
| # ── Vòng synthesis chính (FIX: trước đây main() cụt ở đây, exit 0 giả | |
| # mà không sinh segment nào -> stage TTS DONE giả, mix chạy trên thư mục rỗng) | |
| failed_segments = [] | |
| manifest_entries = [] | |
| ok_count = 0 | |
| skip_count = 0 | |
| generatable = 0 | |
| total_units = len(blocks) | |
| print(f"[TTS] Synthesizing {total_units} unit(s) with engine '{args.engine}' voice '{args.voice}'...") | |
| for pos, unit in enumerate(blocks, 1): | |
| uid = str(unit.get("id")) | |
| start_ms = int(unit.get("start_ms", 0) or 0) | |
| end_ms = int(unit.get("end_ms", start_ms) or start_ms) | |
| raw_text = str(unit.get("text", "") or "") | |
| try: | |
| norm_text = normalizer.normalize(raw_text) if hasattr(normalizer, "normalize") else raw_text | |
| except Exception: | |
| norm_text = raw_text | |
| norm_text = str(norm_text or "").strip() | |
| seg_wav = segments_dir / (f"{int(uid):04d}.wav" if uid.isdigit() else f"seg_{uid}.wav") | |
| if (not _block_should_generate_tts(unit, action_by_id, labels, skipped_tts_ids) | |
| or not norm_text or _is_tts_placeholder(norm_text)): | |
| skip_count += 1 | |
| manifest_entries.append({ | |
| "id": uid, "start_ms": start_ms, "end_ms": end_ms, | |
| "text": raw_text, "segment_wav": None, "status": "SKIPPED", | |
| }) | |
| continue | |
| generatable += 1 | |
| # Tái dùng segment hợp lệ còn sót (plan đổi đã xóa stale ở _prepare_tts_plan) | |
| if seg_wav.exists() and seg_wav.stat().st_size > 500: | |
| ok_count += 1 | |
| manifest_entries.append({ | |
| "id": uid, "start_ms": start_ms, "end_ms": end_ms, | |
| "text": norm_text, "segment_wav": str(seg_wav), "status": "CACHED", | |
| }) | |
| continue | |
| try: | |
| synthesize_text_to_wav(norm_text, seg_wav, args, tts_config, | |
| piper_voice=piper_voice, | |
| omnivoice_session=omnivoice_session) | |
| if not verify_audio_file(seg_wav, ffmpeg_path=str(ffmpeg_path)): | |
| raise RuntimeError("synthesis returned empty/invalid audio") | |
| ok_count += 1 | |
| manifest_entries.append({ | |
| "id": uid, "start_ms": start_ms, "end_ms": end_ms, | |
| "text": norm_text, "segment_wav": str(seg_wav), "status": "OK", | |
| }) | |
| except Exception as e: | |
| err = str(e)[:200] | |
| print(f"[TTS] Unit {uid} failed: {err} -> silence placeholder", file=sys.stderr) | |
| failed_segments.append({"id": uid, "text": norm_text, "reason": err}) | |
| try: | |
| slot_ms = max(300, end_ms - start_ms) | |
| AudioSegment.silent(duration=slot_ms).export(str(seg_wav), format="wav") | |
| except Exception: | |
| pass | |
| manifest_entries.append({ | |
| "id": uid, "start_ms": start_ms, "end_ms": end_ms, | |
| "text": norm_text, "segment_wav": str(seg_wav) if seg_wav.exists() else None, | |
| "status": "FAILED_SILENCE", | |
| }) | |
| if pos % 10 == 0 or pos == total_units: | |
| pct = 10 + int(80 * pos / max(1, total_units)) | |
| print(f"PROGRESS: {pct}%", flush=True) | |
| write_reports(output_wav.parent, failed_segments, manifest_entries) | |
| print(f"[TTS] Done: {ok_count}/{generatable} synthesized, {skip_count} skipped, " | |
| f"{len(failed_segments)} failed->silence. Manifest: {len(manifest_entries)} entries.") | |
| GPUResourceManager.instance().clear_vram_cache() | |
| GPUResourceManager.instance().unload_ollama_model() | |
| sys.exit(0) | |
| if __name__ == "__main__": | |
| main() | |