DRIPPY4 / app /core /tts_worker_cli.py
hoangtaiii's picture
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
Raw History Blame Contribute Delete
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()