thai_indextts2 / webui_parallel_thai.py
williampike's picture
Upload folder using huggingface_hub
4d3248c verified
Raw
History Blame Contribute Delete
114 kB
from __future__ import annotations
import argparse
import atexit
import itertools
import json
import logging
import math
import multiprocessing as mp
import os
import random
import shutil
import sys
import threading
import time
import gc
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Dict, List, Optional
import warnings
import numpy as np
import torch
warnings.filterwarnings("ignore", category=FutureWarning)
warnings.filterwarnings("ignore", category=UserWarning)
import gradio as gr
import pandas as pd
from omegaconf import OmegaConf
current_dir = os.path.dirname(os.path.abspath(__file__))
sys.path.append(current_dir)
sys.path.append(os.path.join(current_dir, "indextts"))
from tools.i18n.i18n import I18nAuto
parser = argparse.ArgumentParser(description="IndexTTS Parallel WebUI")
parser.add_argument("--verbose", action="store_true", default=False, help="Enable verbose logging")
parser.add_argument("--port", type=int, default=7862, help="Port for the web UI")
parser.add_argument("--host", type=str, default="0.0.0.0", help="Host for the web UI")
parser.add_argument("--model_dir", type=str, default="checkpoints", help="Model checkpoints directory")
parser.add_argument("--is_fp16", action="store_true", default=False, help="Enable fp16 inference")
cmd_args = parser.parse_args()
if not os.path.exists(cmd_args.model_dir):
print(f"Model directory {cmd_args.model_dir} does not exist. Please download the model first.")
sys.exit(1)
required_files = [
"config.yaml",
"s2mel.pth",
"wav2vec2bert_stats.pt",
]
for file_name in required_files:
file_path = os.path.join(cmd_args.model_dir, file_name)
if not os.path.exists(file_path):
print(f"Required file {file_path} does not exist. Please download it.")
sys.exit(1)
try:
BASE_CFG = OmegaConf.load(os.path.join(cmd_args.model_dir, "config.yaml"))
except Exception as exc: # pragma: no cover - config must load
print(f"Failed to load config.yaml: {exc}")
sys.exit(1)
hf_cache_dir = os.path.join(cmd_args.model_dir, "hf_cache")
torch_cache_dir = os.path.join(cmd_args.model_dir, "torch_cache")
os.environ.setdefault("INDEXTTS_USE_DEEPSPEED", "0")
os.environ.setdefault("HF_HOME", hf_cache_dir)
os.environ.setdefault("HF_HUB_CACHE", hf_cache_dir)
os.environ.setdefault("TRANSFORMERS_CACHE", hf_cache_dir)
os.environ.setdefault("TORCH_HOME", torch_cache_dir)
os.makedirs(hf_cache_dir, exist_ok=True)
os.makedirs(torch_cache_dir, exist_ok=True)
from indextts.infer_v2_thai import IndexTTS2
from text_preprocessor import ThaiTextPreprocessor
i18n = I18nAuto(language="Auto")
logger = logging.getLogger("webui_parallel")
os.makedirs(os.path.join(current_dir, "outputs", "tasks"), exist_ok=True)
os.makedirs(os.path.join(current_dir, "prompts"), exist_ok=True)
os.environ.setdefault("INDEXTTS_USE_DEEPSPEED", "0")
example_cases: List[List[Any]] = []
examples_path = Path(current_dir) / "examples" / "cases.jsonl"
if examples_path.exists():
with examples_path.open("r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if not line:
continue
example = json.loads(line)
emo_audio = example.get("emo_audio")
emo_audio_path = os.path.join("examples", emo_audio) if emo_audio else None
example_cases.append([
os.path.join("examples", example.get("prompt_audio", "sample_prompt.wav")),
example.get("emo_mode", 0),
example.get("text"),
emo_audio_path,
example.get("emo_weight", 1.0),
example.get("emo_text", ""),
example.get("emo_vec_1", 0),
example.get("emo_vec_2", 0),
example.get("emo_vec_3", 0),
example.get("emo_vec_4", 0),
example.get("emo_vec_5", 0),
])
EMO_CHOICES = [
"Match prompt audio",
"Use emotion reference audio",
"Use emotion vector (Thai 5-Emo)",
"Use emotion text description",
"Use emotion vector (Original 8-Emo)",
]
parallel_worker_config = {
"model_dir": cmd_args.model_dir,
"is_fp16": cmd_args.is_fp16,
"verbose": cmd_args.verbose,
"hf_cache": hf_cache_dir,
"torch_cache": torch_cache_dir,
"gpt_path": None,
"bpe_path": None,
}
class WorkerPool:
def __init__(self, config: Dict[str, Any]):
self.config = config
self.ctx = mp.get_context("spawn")
self.job_queue: Optional[mp.Queue] = None
self.result_queue: Optional[mp.Queue] = None
self.processes: List[mp.Process] = []
self.worker_count = 0
self.lock = threading.Lock()
self.batch_counter = itertools.count()
def _all_alive(self) -> bool:
return all(p.is_alive() for p in self.processes)
def ensure(self, count: int):
count = max(1, int(count))
with self.lock:
if self.worker_count == count and self.processes and self._all_alive():
return
self.stop_locked()
self.start_locked(count)
def start_locked(self, count: int):
self.job_queue = self.ctx.Queue()
self.result_queue = self.ctx.Queue()
self.processes = []
self.worker_count = count
for _ in range(count):
p = self.ctx.Process(
target=_worker_loop,
args=(self.job_queue, self.result_queue, self.config),
daemon=True)
p.start()
self.processes.append(p)
def stop_locked(self):
if not self.processes:
return
if self.job_queue is not None:
for _ in self.processes:
self.job_queue.put({"type": "stop"})
for p in self.processes:
p.join(timeout=5)
self.processes = []
if self.job_queue is not None:
self.job_queue.close()
self.job_queue = None
if self.result_queue is not None:
self.result_queue.close()
self.result_queue = None
self.worker_count = 0
def stop(self):
with self.lock:
self.stop_locked()
def run_jobs(self, jobs: List[GenerationJob], progress: Optional[gr.Progress]):
if not jobs:
return {}
with self.lock:
if not self.processes or self.job_queue is None or self.result_queue is None:
raise RuntimeError("Worker pool not initialized")
batch_id = next(self.batch_counter)
total = len(jobs)
for job in jobs:
payload = job.__dict__.copy()
payload["batch_id"] = batch_id
self.job_queue.put(payload)
row_results: Dict[int, Dict[str, Any]] = {}
processed = 0
total = len(jobs)
while processed < total:
message = self.result_queue.get() # type: ignore[arg-type]
if message.get("type") == "init_error":
raise RuntimeError(f"Worker failed to start: {message['error']}")
if message.get("batch_id") != batch_id:
continue
row_results[message["row_id"]] = message
processed += 1
_update_progress(progress, min(processed / total, 0.999), desc=f"Processed {processed}/{total}")
_update_progress(progress, 1.0, desc="Parallel generation complete")
return row_results
worker_pool = WorkerPool(parallel_worker_config)
def _shutdown_worker_pool():
worker_pool.stop()
atexit.register(_shutdown_worker_pool)
_PRIMARY_TTS: Optional[IndexTTS2] = None
_MODEL_SELECTION: Dict[str, Optional[str]] = {
"gpt": r"C:\datasetmaker\index-tts\models\thaiseperate2.pth",
"bpe": r"C:\datasetmaker\index-tts\checkpoints\thai_segmented_bpe.model"
}
def _candidate_paths(base_dirs: List[Path], suffixes: List[str]) -> List[str]:
results: List[str] = []
seen: set[str] = set()
for base in base_dirs:
if not base or not base.exists():
continue
for suffix in suffixes:
for path in base.glob(f"*{suffix}"):
resolved = str(path.resolve())
if resolved not in seen:
seen.add(resolved)
results.append(resolved)
results.sort()
return results
def _is_gpt_checkpoint(path: Path) -> bool:
name = path.name.lower()
if not name.endswith(".pth"):
return False
excluded = ("s2mel", "campplus", "bigvgan", "wav2vec", "emo", "spk", "cfm")
return not any(token in name for token in excluded)
def _discover_gpt_checkpoints() -> List[str]:
bases = [
Path(cmd_args.model_dir),
Path(current_dir) / "models",
]
candidates = _candidate_paths(bases, [".pth"])
return [path for path in candidates if _is_gpt_checkpoint(Path(path))]
def _discover_bpe_models() -> List[str]:
bases = [
Path(cmd_args.model_dir),
Path(current_dir) / "tokenizers",
]
return _candidate_paths(bases, [".model"])
def dispose_primary_tts():
global _PRIMARY_TTS
if _PRIMARY_TTS is not None:
try:
if hasattr(_PRIMARY_TTS, "gr_progress"):
_PRIMARY_TTS.gr_progress = None
finally:
_PRIMARY_TTS = None
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
def build_primary_tts() -> IndexTTS2:
if _MODEL_SELECTION["gpt"] is None or _MODEL_SELECTION["bpe"] is None:
raise RuntimeError("Model selection is not set. Provide GPT and BPE paths before loading.")
return IndexTTS2(
model_dir=cmd_args.model_dir,
cfg_path=os.path.join(cmd_args.model_dir, "config.yaml"),
is_fp16=cmd_args.is_fp16,
use_cuda_kernel=False,
use_accel=True,
use_torch_compile=False,
gpt_checkpoint_path=_MODEL_SELECTION["gpt"],
bpe_model_path=_MODEL_SELECTION["bpe"])
def load_primary_tts(gpt_path: str, bpe_path: str) -> IndexTTS2:
dispose_primary_tts()
resolved_gpt = os.path.abspath(gpt_path)
resolved_bpe = os.path.abspath(bpe_path)
previous_selection = _MODEL_SELECTION.copy()
_MODEL_SELECTION["gpt"] = resolved_gpt
_MODEL_SELECTION["bpe"] = resolved_bpe
try:
tts = build_primary_tts()
except Exception:
_MODEL_SELECTION.update(previous_selection)
dispose_primary_tts()
raise
global _PRIMARY_TTS
_PRIMARY_TTS = tts
parallel_worker_config["gpt_path"] = resolved_gpt
parallel_worker_config["bpe_path"] = resolved_bpe
worker_pool.stop()
return tts
def ensure_primary_tts() -> IndexTTS2:
if _PRIMARY_TTS is None:
raise RuntimeError("No GPT checkpoint loaded. Use the Load button in the UI.")
return _PRIMARY_TTS
def _model_status_text() -> str:
if _PRIMARY_TTS is None:
return "⚠️ No model loaded. Select a GPT checkpoint and BPE tokenizer, then click Load."
gpt_path = _MODEL_SELECTION.get("gpt")
bpe_path = _MODEL_SELECTION.get("bpe")
gpt_name = Path(gpt_path).name if gpt_path else "?"
bpe_name = Path(bpe_path).name if bpe_path else "?"
return f"✅ Loaded GPT: **{gpt_name}** | BPE: **{bpe_name}**"
def _format_label(path: str) -> str:
path_obj = Path(path)
candidates: List[str] = []
try:
rel_model = os.path.relpath(path, cmd_args.model_dir)
if not rel_model.startswith(".."):
prefix = Path(cmd_args.model_dir).name or "checkpoints"
candidates.append(f"{prefix}/{rel_model}".replace("\\", "/"))
except ValueError:
pass
try:
rel_repo = os.path.relpath(path, current_dir)
if not rel_repo.startswith(".."):
candidates.append(rel_repo.replace("\\", "/"))
except ValueError:
pass
candidates.append(path_obj.name)
for label in candidates:
if label:
return label
return str(path_obj)
def _format_dropdown_choices(
paths: List[str],
current_selection: Optional[str]) -> Tuple[List[str], Dict[str, str], Optional[str]]:
labels: List[str] = []
mapping: Dict[str, str] = {}
selected_label: Optional[str] = None
for path in paths:
label = _format_label(path)
base_label = label
suffix = 1
while label in mapping:
label = f"{base_label} ({suffix})"
suffix += 1
mapping[label] = path
labels.append(label)
if current_selection and os.path.abspath(path) == os.path.abspath(current_selection):
selected_label = label
if labels and selected_label is None:
selected_label = labels[0]
return labels, mapping, selected_label
@dataclass
class GenerationJob:
row_id: int
prompt_path: str
text: str
output_path: str
emo_mode: int
emo_weight: float
emo_vector: Optional[List[float]]
emo_text: str
emo_random: bool
emo_ref_path: Optional[str]
max_tokens: int
generation_kwargs: Dict[str, Any]
verbose: bool
duration_seconds: Optional[float] = None
accent_ref_path: Optional[str] = None
def _normalize_seed(seed_value: Any) -> Optional[int]:
if seed_value is None:
return None
if isinstance(seed_value, str):
value = seed_value.strip()
if not value:
return None
try:
seed = int(value)
except ValueError:
try:
seed = int(float(value))
except ValueError:
return None
elif isinstance(seed_value, bool):
seed = int(seed_value)
elif isinstance(seed_value, float):
if math.isnan(seed_value):
return None
seed = int(seed_value)
else:
try:
seed = int(seed_value)
except (TypeError, ValueError):
return None
if seed < 0:
seed = abs(seed)
return seed
def _normalize_duration_seconds(value: Any) -> Optional[float]:
if value is None:
return None
if isinstance(value, str):
value = value.strip()
if not value:
return None
try:
seconds = float(value)
except (TypeError, ValueError):
return None
if seconds <= 0:
return None
return seconds
def _apply_seed(seed: Optional[int]) -> None:
if seed is None:
return
py_seed = int(seed % (2**32))
random.seed(py_seed)
np.random.seed(py_seed)
torch_seed = int(seed % (2**63 - 1))
torch.manual_seed(torch_seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(torch_seed)
def _prepare_generation_kwargs(raw_kwargs: Dict[str, Any]) -> Dict[str, Any]:
kwargs = dict(raw_kwargs or {})
seed = _normalize_seed(kwargs.pop("seed", None))
_apply_seed(seed)
return kwargs
def trim_audio_silences(path: str, max_sec: float = 1.0) -> str:
try:
import librosa
import soundfile as sf
import numpy as np
y, sr = librosa.load(path, sr=None)
# 1. ค้นหาช่วงที่ไม่ใช่เสียงเงียบ (top_db=28 เพื่อตัด noise ที่เบากว่าเสียงพูดทิ้ง)
intervals = librosa.effects.split(y, top_db=28, frame_length=2048, hop_length=512)
if len(intervals) == 0:
return path
pieces = []
max_pad = int(max_sec * sr)
# 2. จัดการเสียงเงียบ: ตัดหัวท้ายทิ้ง 100% และคุมจังหวะเงียบตรงกลางไม่ให้เกิน 1 วินาที
for i, intv in enumerate(intervals):
# เพิ่มช่วงที่มีเสียง
pieces.append(y[intv[0]:intv[1]])
# ถ้ามีช่วงถัดไป ให้เช็คช่วงเงียบตรงกลาง
if i < len(intervals) - 1:
gap_len = intervals[i+1][0] - intv[1]
if gap_len > max_pad:
# ถ้าเงียบเกิน 1 วิ ให้เหลือแค่ 1 วิ
pieces.append(np.zeros(max_pad, dtype=y.dtype))
elif gap_len > 0:
# ถ้าเงียบไม่เกิน 1 วิ ให้คงไว้ตามธรรมชาติ
pieces.append(y[intv[1]:intervals[i+1][0]])
y_out = np.concatenate(pieces)
sf.write(path, y_out, sr)
except Exception as e:
print("Trim silence error:", e)
return path
def _worker_loop(job_queue: mp.Queue, result_queue: mp.Queue, config: Dict[str, Any]):
hf_cache = config.get("hf_cache")
torch_cache = config.get("torch_cache")
if hf_cache:
os.environ.setdefault("HF_HOME", hf_cache)
os.environ.setdefault("HF_HUB_CACHE", hf_cache)
os.environ.setdefault("TRANSFORMERS_CACHE", hf_cache)
os.makedirs(hf_cache, exist_ok=True)
if torch_cache:
os.environ.setdefault("TORCH_HOME", torch_cache)
os.makedirs(torch_cache, exist_ok=True)
os.environ.setdefault("INDEXTTS_USE_DEEPSPEED", "0")
gpt_override = config.get("gpt_path")
bpe_override = config.get("bpe_path")
if not gpt_override or not bpe_override:
result_queue.put({"type": "init_error", "error": "No GPT/BPE model loaded. Use the Load button."})
return
try:
worker_tts = IndexTTS2(
model_dir=config["model_dir"],
cfg_path=os.path.join(config["model_dir"], "config.yaml"),
is_fp16=config.get("is_fp16", False),
use_cuda_kernel=False,
use_accel=True,
use_torch_compile=False,
gpt_checkpoint_path=gpt_override,
bpe_model_path=bpe_override)
except Exception as exc: # pragma: no cover - worker init path
logger.exception("Worker failed to initialize")
result_queue.put({"type": "init_error", "error": str(exc)})
return
while True:
job = job_queue.get()
if isinstance(job, dict) and job.get("type") == "stop":
break
try:
emo_mode = job["emo_mode"]
emo_audio_prompt = job["emo_ref_path"] if emo_mode == 1 else None
emo_alpha = job["emo_weight"] if emo_mode == 1 else 1.0
emo_vector = job["emo_vector"] if emo_mode == 2 else None
use_emo_text = emo_mode == 3
generation_kwargs = _prepare_generation_kwargs(job.get("generation_kwargs", {}))
trim_silence_value = generation_kwargs.pop("trim_silence", False)
auto_retry_value = generation_kwargs.pop("auto_retry", False)
use_dataset_spacing_value = generation_kwargs.pop("use_dataset_spacing", False)
use_g2p_value = generation_kwargs.pop("use_g2p", False)
preprocessor = ThaiTextPreprocessor(use_g2p=use_g2p_value, use_dataset_spacing=use_dataset_spacing_value)
clean_text = preprocessor.process(job["text"])
prompt_path = job["prompt_path"]
if trim_silence_value and prompt_path and os.path.exists(prompt_path):
trim_audio_silences(prompt_path)
max_retries = 3 if auto_retry_value else 1
for attempt in range(max_retries):
worker_tts.infer(
spk_audio_prompt=prompt_path,
text=clean_text,
output_path=job["output_path"],
emo_audio_prompt=emo_audio_prompt,
emo_alpha=emo_alpha,
emo_vector=emo_vector,
use_emo_text=use_emo_text,
emo_text=job["emo_text"],
use_random=job["emo_random"],
verbose=job.get("verbose", False),
max_text_tokens_per_segment=job["max_tokens"],
duration_seconds=job.get("duration_seconds"),
accent_audio_prompt=job.get("accent_ref_path"),
**generation_kwargs)
if trim_silence_value and os.path.exists(job["output_path"]):
trim_audio_silences(job["output_path"])
if auto_retry_value and os.path.exists(job["output_path"]):
try:
import librosa
y_out, sr_out = librosa.load(job["output_path"], sr=None)
dur = len(y_out) / sr_out
toks = len(worker_tts.tokenizer.tokenize(clean_text))
speed = float(generation_kwargs.get("speed_factor", 1.0))
est = toks * 0.3 * (1.0 / speed)
if (dur < est * 0.4 or dur > est * 2.5) and toks > 5:
if attempt < max_retries - 1:
print(f"⚠️ Worker: Audio length anomaly detected (Dur: {dur:.2f}s, Est: {est:.2f}s). Retrying ({attempt+1}/3)...")
continue
except Exception as e:
print("Retry check error:", e)
break
result_queue.put(
{
"type": "result",
"row_id": job["row_id"],
"status": "Completed",
"output_path": job["output_path"],
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
"error": None,
"batch_id": job.get("batch_id"),
}
)
except Exception as exc: # pragma: no cover - worker runtime path
logger.exception("Worker generation error")
result_queue.put(
{
"type": "result",
"row_id": job["row_id"],
"status": f"Error: {exc}",
"output_path": None,
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
"error": str(exc),
"batch_id": job.get("batch_id"),
}
)
try:
worker_tts.unload() # type: ignore[attr-defined]
except Exception: # pragma: no cover - optional cleanup
pass
def _update_progress(progress: Optional[gr.Progress], value: float, desc: str = "") -> None:
if progress is None:
return
try:
progress(value, desc=desc)
except Exception:
pass
MAX_LENGTH_TO_USE_SPEED = 70
# Web Audio API streaming — signal via gr.HTML data attribute.
#
# Why HTML data attribute instead of Textbox DOM polling:
# - gr.HTML renders raw content we fully control — no DOM structure uncertainty
# - data-v attribute on a <span> is trivially readable: el.getAttribute('data-v')
# - No textarea/input querySelector needed, no container=False ambiguity
#
# Flow:
# Python yields new HTML string (with updated data-v) → Gradio updates innerHTML
# → JS reads data-v from #sph-sig → plays chunk via Web Audio API
def create_demo() -> gr.Blocks:
gpt_choices = _discover_gpt_checkpoints()
bpe_choices = _discover_bpe_models()
gpt_labels, gpt_map, initial_gpt_label = _format_dropdown_choices(gpt_choices, _MODEL_SELECTION["gpt"])
bpe_labels, bpe_map, initial_bpe_label = _format_dropdown_choices(bpe_choices, _MODEL_SELECTION["bpe"])
gpt_cfg = getattr(BASE_CFG, "gpt", {})
max_mel_tokens_limit = int(getattr(gpt_cfg, "max_mel_tokens", 2048))
if max_mel_tokens_limit < 100:
max_mel_tokens_limit = 100
default_mel_value = min(1500, max_mel_tokens_limit)
max_text_tokens_limit = int(getattr(gpt_cfg, "max_text_tokens", 256))
if max_text_tokens_limit < 40:
max_text_tokens_limit = 40
default_text_tokens = min(120, max_text_tokens_limit)
cfg_version = getattr(BASE_CFG, "version", "1.0")
outputs_dir = os.path.join(current_dir, "outputs")
os.makedirs(outputs_dir, exist_ok=True)
with gr.Blocks(title="IndexTTS Parallel Demo") as demo:
model_status = gr.Markdown(value=_model_status_text())
gpt_map_state = gr.State(gpt_map)
bpe_map_state = gr.State(bpe_map)
with gr.Row():
gpt_dropdown = gr.Dropdown(
choices=gpt_labels,
value=initial_gpt_label,
label="GPT Checkpoint (.pth)",
interactive=True)
bpe_dropdown = gr.Dropdown(
choices=bpe_labels,
value=initial_bpe_label,
label="BPE Tokenizer (.model)",
interactive=True)
refresh_models_button = gr.Button("Refresh Models", variant="secondary")
load_models_button = gr.Button("Load Models", variant="primary")
def refresh_model_lists():
gpt_files = _discover_gpt_checkpoints()
bpe_files = _discover_bpe_models()
gpt_labels_new, gpt_map_new, gpt_value = _format_dropdown_choices(gpt_files, _MODEL_SELECTION["gpt"])
bpe_labels_new, bpe_map_new, bpe_value = _format_dropdown_choices(bpe_files, _MODEL_SELECTION["bpe"])
return (
gr.update(choices=gpt_labels_new, value=gpt_value),
gr.update(choices=bpe_labels_new, value=bpe_value),
gpt_map_new,
bpe_map_new,
_model_status_text())
def handle_model_load(
gpt_label: Optional[str],
bpe_label: Optional[str],
gpt_map_value: Optional[Dict[str, str]],
bpe_map_value: Optional[Dict[str, str]],
progress: gr.Progress = gr.Progress(track_tqdm=False)) -> str:
gpt_map_local = gpt_map_value or {}
bpe_map_local = bpe_map_value or {}
gpt_path = gpt_map_local.get(gpt_label or "", gpt_label)
bpe_path = bpe_map_local.get(bpe_label or "", bpe_label)
if not gpt_path or not bpe_path:
gr.Warning("Select both a GPT checkpoint and a BPE tokenizer before loading.")
return _model_status_text()
progress(0.1, "Loading models...")
try:
load_primary_tts(gpt_path, bpe_path)
except Exception as exc:
logger.exception("Failed to load models")
gr.Warning(f"Failed to load models: {exc}")
return f"❌ Failed to load models: {exc}"
gr.Info("Models loaded successfully.")
return _model_status_text()
refresh_models_button.click(
refresh_model_lists,
inputs=[],
outputs=[gpt_dropdown, bpe_dropdown, gpt_map_state, bpe_map_state, model_status])
load_models_button.click(
handle_model_load,
inputs=[gpt_dropdown, bpe_dropdown, gpt_map_state, bpe_map_state],
outputs=model_status)
batch_rows_state = gr.State([])
next_batch_id_state = gr.State(1)
gr.HTML(
"""
<h2 style=\"text-align:center;\">IndexTTS2 Parallel Batch Demo</h2>
"""
)
with gr.Accordion("Emotion Settings", open=True):
with gr.Row():
emo_control_method = gr.Radio(
choices=EMO_CHOICES,
type="index",
value=0,
label="Emotion Control Mode")
with gr.Group(visible=True) as emo_weight_group:
with gr.Row():
emo_weight = gr.Slider(label="Emotion Weight", minimum=0.0, maximum=1.6, value=0.8, step=0.01)
with gr.Group(visible=False) as emotion_reference_group:
with gr.Row():
emo_upload = gr.Audio(label="Emotion Reference Audio", type="filepath")
with gr.Row():
emo_random = gr.Checkbox(label="Random Emotion Sampling", value=False, visible=False)
with gr.Group(visible=False) as thai_emotion_vector_group:
with gr.Row():
with gr.Column():
tvec1 = gr.Slider(label="Neutral", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
tvec2 = gr.Slider(label="Angry", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
tvec3 = gr.Slider(label="Happy", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
with gr.Column():
tvec4 = gr.Slider(label="Sad", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
tvec5 = gr.Slider(label="Frustrated", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
with gr.Group(visible=False) as emotion_vector_group:
with gr.Row():
with gr.Column():
vec1 = gr.Slider(label="Joy", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
vec2 = gr.Slider(label="Anger", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
vec3 = gr.Slider(label="Sadness", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
vec4 = gr.Slider(label="Fear", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
with gr.Column():
vec5 = gr.Slider(label="Disgust", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
vec6 = gr.Slider(label="Low Mood", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
vec7 = gr.Slider(label="Surprise", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
vec8 = gr.Slider(label="Calm", minimum=0.0, maximum=1.4, value=0.0, step=0.05)
with gr.Group(visible=False) as emo_text_group:
emo_text = gr.Textbox(label="Emotion Description", placeholder="Describe the target emotion", value="")
with gr.Accordion("Advanced Generation Settings", open=False):
with gr.Row():
with gr.Column(scale=1):
gr.Markdown("**GPT2 Sampling Settings**")
with gr.Row():
do_sample = gr.Checkbox(label="do_sample", value=True, info="Enable sampling")
temperature = gr.Slider(label="temperature", minimum=0.1, maximum=2.0, value=0.8, step=0.1)
with gr.Row():
top_p = gr.Slider(label="top_p", minimum=0.0, maximum=1.0, value=0.8, step=0.01)
top_k = gr.Slider(label="top_k", minimum=0, maximum=100, value=30, step=1)
num_beams = gr.Slider(label="num_beams", value=3, minimum=1, maximum=10, step=1)
with gr.Row():
repetition_penalty = gr.Number(label="repetition_penalty", precision=None, value=10.0, minimum=0.1, maximum=20.0, step=0.1)
length_penalty = gr.Number(label="length_penalty", precision=None, value=0.0, minimum=-2.0, maximum=2.0, step=0.1)
max_mel_tokens = gr.Slider(
label="max_mel_tokens",
value=default_mel_value,
minimum=50,
maximum=max_mel_tokens_limit,
step=10,
info="Maximum generated mel tokens")
seed_value = gr.Number(
label="Seed",
value=None,
precision=0,
minimum=0,
step=1,
info="Leave blank for random sampling; set a value for reproducible outputs.")
gr.Markdown("**Voice & Timing Settings**")
speed_factor = gr.Slider(
label="Speed Rate (ความเร็ว: < 1 เร็ว, > 1 ช้า)",
minimum=0.5,
maximum=2.0,
value=1.0,
step=0.1,
info="ปรับความเร็วการพูดของ AI")
interval_silence = gr.Slider(
label="Interval Silence (ms)",
minimum=0,
maximum=1000,
value=200,
step=50,
info="ระยะเวลาพักหายใจระหว่างประโยค")
use_g2p = gr.Checkbox(label="🪄 โหมดสะกดคำง่าย (G2P)", value=False, info="แปลงคำยากๆ ให้สะกดตรงตัวก่อนพากย์ (เช่น สุทธิกร -> สุดทิกอน)")
use_dataset_spacing = gr.Checkbox(label="✂️ แบ่งคำและจัด Spacebar แบบ Dataset", value=False, info="ประมวลผลข้อความให้มีการเว้นวรรค 1-2 ช่อง เพื่อให้ตรงกับโมเดล BPE ตัวใหม่")
classic_mode = gr.Checkbox(
label="✅ Classic Mode (โหมดดั้งเดิม)",
value=False,
info="ติ๊กเพื่อข้ามระบบแยกสำเนียง/ความเร็ว แล้วรันด้วยลอจิกดั้งเดิม"
)
trim_silence = gr.Checkbox(label="✂️ Trim Silence (ตัดเสียงเงียบลากยาว)", value=False, info="ถ้าผลลัพธ์หรือเสียงต้นฉบับมีช่วงเงียบเกิน 1 วินาที จะตัดให้เหลือแค่ 1 วินาที")
auto_retry = gr.Checkbox(label="🔁 Auto-Regenerate (ป้องกันอาการเอ๋อ)", value=False, info="ถ้า AI สร้างเสียงยาวเกินไปหรือสั้นผิดปกติเมื่อเทียบกับจำนวนคำ จะสั่ง Gen ใหม่ให้อัตโนมัติ")
chain_segments = gr.Checkbox(label="🔗 Chain Segments (คงอารมณ์เสียงให้ต่อเนื่อง)", value=False, info="เมื่อพิมพ์ข้อความยาวจนโดนหั่นเป็น 2 ท่อน จะดึงเสียงท่อนแรกมาเป็นต้นแบบให้ท่อนต่อไปเสมอ (อารมณ์/เสียงไม่แกว่ง)")
dur_per_token = gr.Slider(label="⏱️ Auto-Regen Sensitivity (Duration/Token)", value=0.12, minimum=0.05, maximum=0.5, step=0.01, info="ค่าเฉลี่ยความยาววินาทีต่อ 1 Token (ถ้าเสียงที่ Gen ได้สั้นหรือยาวกว่าค่านี้มากๆ ระบบจะ Gen ใหม่)")
with gr.Column(scale=2):
gr.Markdown("**Sentence Settings**")
max_text_tokens_per_sentence = gr.Slider(
label="Max tokens per sentence",
value=default_text_tokens,
minimum=20,
maximum=max_text_tokens_limit,
step=2,
key="max_text_tokens_per_sentence")
duration_seconds_input = gr.Number(
label="Target duration (seconds)",
value=None,
precision=2,
minimum=0,
step=0.1,
info="Optional: approximate overall audio length. Leave blank for free duration.")
with gr.Accordion("Preview sentences", open=True):
sentences_preview = gr.Dataframe(
headers=["Index", "Sentence", "Token Count"],
key="sentences_preview",
wrap=True)
# [FIX] นำ use_g2p เข้ามาอยู่ในกลุ่ม advanced_params เพื่อการแยกตัวแปรที่สมบูรณ์!
advanced_params = [
do_sample,
top_p,
top_k,
temperature,
length_penalty,
num_beams,
repetition_penalty,
max_mel_tokens,
seed_value,
speed_factor,
interval_silence,
classic_mode,
use_g2p,
use_dataset_spacing,
trim_silence,
auto_retry,
chain_segments,
dur_per_token,
]
def build_generation_kwargs(
do_sample_value,
top_p_value,
top_k_value,
temperature_value,
length_penalty_value,
num_beams_value,
repetition_penalty_value,
max_mel_tokens_value,
seed_value,
speed_factor_value,
interval_silence_value,
classic_mode_value,
use_g2p_value,
use_dataset_spacing_value=False,
trim_silence_value=False,
auto_retry_value=False,
chain_segments_value=False,
dur_per_token_value=0.12
):
try:
top_k_int = int(top_k_value)
except (TypeError, ValueError):
top_k_int = 0
try:
num_beams_int = int(num_beams_value)
except (TypeError, ValueError):
num_beams_int = 1
kwargs = {
"do_sample": bool(do_sample_value),
"top_p": float(top_p_value),
"top_k": top_k_int if top_k_int > 0 else None,
"temperature": float(temperature_value),
"length_penalty": float(length_penalty_value),
"num_beams": num_beams_int,
"repetition_penalty": float(repetition_penalty_value),
"max_mel_tokens": int(max_mel_tokens_value),
"speed_factor": float(speed_factor_value),
"interval_silence": int(interval_silence_value),
"classic_mode": bool(classic_mode_value),
"use_g2p": bool(use_g2p_value),
"use_dataset_spacing": bool(use_dataset_spacing_value),
"trim_silence": bool(trim_silence_value),
"auto_retry": bool(auto_retry_value),
"chain_segments": bool(chain_segments_value),
"dur_per_token": float(dur_per_token_value)
}
seed_int = _normalize_seed(seed_value)
if seed_int is not None:
kwargs["seed"] = seed_int
return kwargs
with gr.Tab("Single Generation"):
with gr.Row():
with gr.Column():
prompt_audio = gr.Audio(label="Voice Reference (เสียงหลักที่ต้องการโคลน)", key="prompt_audio", sources=["upload", "microphone"], type="filepath")
accent_audio = gr.Audio(label="Accent Reference (เสียงคนไทยเพื่อแก้สำเนียง - Optional)", key="accent_audio", sources=["upload", "microphone"], type="filepath")
with gr.Column():
input_text_single = gr.TextArea(
label="Text",
key="input_text_single",
placeholder="Enter text to synthesize",
info=f"Model version {cfg_version}")
with gr.Row():
format_single_btn = gr.Button("🪄 จัดข้อความ (แยกคำ + Spacebar)", variant="secondary")
gen_button = gr.Button("Generate", key="gen_button", interactive=True, variant="primary")
output_audio = gr.Audio(
label="Generated Result (Normal)",
visible=True,
key="output_audio",
autoplay=True
)
stream_audio_output = gr.Audio(
label="Streaming Player (Plays instantly)",
visible=True,
autoplay=True,
streaming=True
)
with gr.Row():
gen_stream_button = gr.Button("Streaming Generate (ทยอย Gen ทีละประโยค)", key="gen_stream_button", interactive=True, variant="secondary")
with gr.Tab("Interactive Segment Builder"):
gr.Markdown("สร้างเสียงทีละท่อน (Segment) เพื่อให้คุณสามารถตรวจสอบและ Regenerate ท่อนที่ไม่พอใจได้ก่อนจะรวมไฟล์")
with gr.Row():
with gr.Column():
seg_prompt_audio = gr.Audio(label="Voice Reference (เสียงหลักที่ต้องการโคลน)", key="seg_prompt_audio", sources=["upload", "microphone"], type="filepath")
seg_accent_audio = gr.Audio(label="Accent Reference (เสียงคนไทยเพื่อแก้สำเนียง - Optional)", key="seg_accent_audio", sources=["upload", "microphone"], type="filepath")
seg_input_text = gr.TextArea(
label="Text",
key="seg_input_text",
placeholder="Enter text to synthesize",
info="ใส่ข้อความทั้งหมด ระบบจะแยกเป็นประโยคให้")
with gr.Row():
seg_format_btn = gr.Button("🪄 จัดข้อความ (แยกคำ + Spacebar)", variant="secondary")
seg_split_btn = gr.Button("1. Split into Segments (แบ่งประโยค)", variant="primary")
seg_status = gr.Markdown("ยังไม่ได้แบ่งประโยค")
with gr.Column():
seg_table = gr.Dataframe(
headers=["Index", "Text", "Status", "Duration (s)"],
datatype=["number", "str", "str", "number"],
interactive=True,
wrap=True)
gr.Markdown("*💡 คลิกที่แต่ละแถวบนตารางด้านบน เพื่อฟังเสียงท่อนนั้นซ้ำ (สำหรับ Check เสียงเฉพาะท่อน)*")
seg_playback = gr.Audio(label="Playback Selected Segment", interactive=False)
with gr.Row():
seg_gen_next_btn = gr.Button("2. Generate Next Segment (สร้างท่อนถัดไป)", variant="primary", interactive=False)
seg_regen_last_btn = gr.Button("Regenerate Last Segment (สร้างท่อนล่าสุดใหม่)", variant="secondary", interactive=False)
seg_clear_btn = gr.Button("Clear All", variant="stop")
current_seg_audio = gr.Audio(label="Current Segment (ท่อนล่าสุด)", interactive=False)
final_seg_audio = gr.Audio(label="Combined Audio (รวมทั้งหมด)", interactive=False)
# Hidden states for Segment Builder
seg_state_texts = gr.State([])
seg_state_wavs = gr.State([])
seg_state_idx = gr.State(0)
with gr.Tab("Batch Generation"):
gr.Markdown("Manage multiple prompt audios, give each its own text, generate in bulk, and retry specific entries as needed.")
with gr.Row():
with gr.Column(scale=2):
with gr.Row():
dataset_path_input = gr.Textbox(
label="Dataset train.txt path",
value="vivy_va_dataset/train.txt",
scale=3,
placeholder="Path to train.txt")
load_dataset_button = gr.Button("Load Dataset", scale=1)
batch_file_input = gr.Files(
label="Add prompt audio files",
file_types=["audio"],
file_count="multiple",
type="filepath")
batch_accent_input = gr.Audio(label="Global Accent Reference for Batch (Optional)", type="filepath")
worker_count = gr.Slider(
label="Parallel workers",
minimum=1,
maximum=8,
value=2,
step=1,
info="Number of parallel TTS workers")
batch_table = gr.Dataframe(
headers=["ID", "Prompt", "Text", "Output", "Status", "Last Generated"],
datatype=["number", "str", "str", "str", "str", "str"],
row_count=(0, "dynamic"),
col_count=6,
interactive=False,
value=[])
with gr.Column():
selected_entry = gr.Dropdown(label="Select entry", choices=[], value=None, interactive=True)
batch_prompt_player = gr.Audio(label="Prompt Audio", type="filepath", interactive=False)
batch_output_player = gr.Audio(label="Generated Audio", type="filepath", interactive=False)
batch_text_input = gr.TextArea(label="Text", placeholder="Enter text for this entry", interactive=True)
with gr.Row():
format_batch_btn = gr.Button("🪄 จัดข้อความ (แยกคำ + Spacebar)", variant="secondary")
apply_text_button = gr.Button("Save Text", variant="primary")
batch_status = gr.Markdown(value="No entry selected.")
with gr.Row():
generate_all_button = gr.Button("Generate All")
regenerate_button = gr.Button("Regenerate Selected")
with gr.Row():
delete_entry_button = gr.Button("Delete Selected")
clear_entries_button = gr.Button("Clear All")
def gen_single(
emo_control_method_value,
prompt,
accent_ref_path,
text,
emo_ref_path,
emo_weight_value,
tvec1_value, tvec2_value, tvec3_value, tvec4_value, tvec5_value,
vec1_value, vec2_value, vec3_value, vec4_value, vec5_value, vec6_value, vec7_value, vec8_value,
emo_text_value,
emo_random_value,
max_text_tokens_per_sentence_value,
duration_seconds_value,
*args,
progress: gr.Progress = gr.Progress()):
if not prompt:
gr.Warning("Upload a prompt audio file first.")
yield gr.update()
return
output_path = os.path.join(current_dir, "outputs", f"spk_{int(time.time())}.wav")
try:
tts = ensure_primary_tts()
except RuntimeError as exc:
gr.Warning(str(exc))
yield gr.update()
return
tts.gr_progress = progress
advanced_values = list(args)
expected_len = len(advanced_params)
if len(advanced_values) < expected_len:
advanced_values.extend([None] * (expected_len - len(advanced_values)))
raw_generation_kwargs = build_generation_kwargs(*advanced_values[:expected_len])
use_g2p_value = raw_generation_kwargs.pop("use_g2p", False)
use_dataset_spacing_value = raw_generation_kwargs.pop("use_dataset_spacing", False)
trim_silence_value = raw_generation_kwargs.pop("trim_silence", False)
auto_retry_value = raw_generation_kwargs.pop("auto_retry", False)
dur_per_token_value = raw_generation_kwargs.pop("dur_per_token", 0.12)
# Re-pop chain_segments to avoid passing it to infer
chain_segments_value = raw_generation_kwargs.pop("chain_segments", False)
generation_kwargs = _prepare_generation_kwargs(raw_generation_kwargs)
# Add chain_segments to generation_kwargs for infer_generator
generation_kwargs["chain_segments"] = chain_segments_value
emo_mode = emo_control_method_value if isinstance(emo_control_method_value, int) else getattr(emo_control_method_value, "value", 0)
tvec_values = [tvec1_value, tvec2_value, tvec3_value, tvec4_value, tvec5_value]
vec_values = [vec1_value, vec2_value, vec3_value, vec4_value, vec5_value, vec6_value, vec7_value, vec8_value]
if emo_mode == 2:
if sum(tvec_values) > 1.5:
gr.Warning("Thai Emotion vector sum cannot exceed 1.5. Adjust the sliders and retry.")
yield gr.update()
return
emo_vector = tvec_values
elif emo_mode == 4:
if sum(vec_values) > 1.5:
gr.Warning("Original Emotion vector sum cannot exceed 1.5. Adjust the sliders and retry.")
yield gr.update()
return
emo_vector = vec_values
else:
emo_vector = None
duration_seconds = _normalize_duration_seconds(duration_seconds_value)
# --- เริ่มกระบวนการแปลงข้อความ ---
preprocessor = ThaiTextPreprocessor(use_g2p=use_g2p_value, use_dataset_spacing=use_dataset_spacing_value)
clean_text = preprocessor.process(text)
print(f"📝 Original Text : {text}")
print(f"✨ Cleaned Text : {clean_text}")
# -------------------------------
if trim_silence_value and prompt and os.path.exists(prompt):
trim_audio_silences(prompt)
max_retries = 3 if auto_retry_value else 1
for attempt in range(max_retries):
try:
tts.infer(
spk_audio_prompt=prompt,
text=clean_text,
output_path=output_path,
emo_audio_prompt=emo_ref_path if emo_mode == 1 else None,
emo_alpha=float(emo_weight_value) if emo_mode in (0, 1) else 1.0,
emo_vector=emo_vector if emo_mode in (2, 4) else None,
use_emo_text=(emo_mode == 3),
emo_text=emo_text_value,
use_random=emo_random_value,
verbose=cmd_args.verbose,
max_text_tokens_per_segment=int(max_text_tokens_per_sentence_value),
duration_seconds=duration_seconds,
accent_audio_prompt=accent_ref_path,
**generation_kwargs)
except AssertionError:
gr.Warning(
"Text segment is too long for the tokenizer with the current "
"'Max tokens per sentence' setting. Try reducing it or splitting "
"the text into shorter sentences.")
yield gr.update()
return
if trim_silence_value and os.path.exists(output_path):
trim_audio_silences(output_path)
if auto_retry_value and os.path.exists(output_path):
try:
import librosa
y_out, sr_out = librosa.load(output_path, sr=None)
dur = len(y_out) / sr_out
toks = len(tts.tokenizer.tokenize(clean_text))
speed = float(generation_kwargs.get("speed_factor", 1.0))
# Log stats to JSONL for future calculation refinement
try:
log_dir = os.path.join(current_dir, "omniman2")
os.makedirs(log_dir, exist_ok=True)
log_file = os.path.join(log_dir, "generation_stats.jsonl")
with open(log_file, "a", encoding="utf-8") as f:
log_entry = {
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
"tokens": toks,
"duration": round(dur, 3),
"dur_per_token": round(dur / toks, 4) if toks > 0 else 0,
"speed_factor": speed,
"text_snippet": clean_text[:100]
}
f.write(json.dumps(log_entry, ensure_ascii=False) + "\n")
except Exception as log_err:
print(f"Log error: {log_err}")
# ------------------------------------------------------
# PRO REGEN LOGIC: ใช้เกณฑ์ความแม่นยำสูง (0.08 - 0.21 s/token)
# ------------------------------------------------------
speed = float(generation_kwargs.get("speed_factor", 1.0))
# ปรับเกณฑ์ตาม Speed (ถ้าปรับสปีด 2x เกณฑ์ก็ต้องหาร 2)
min_limit = 0.08 * (1.0 / speed)
max_limit = 0.21 * (1.0 / speed)
dur_per_tok = dur / toks if toks > 0 else 0
# Anomaly Detection
is_anomaly = (dur_per_tok < min_limit or dur_per_tok > max_limit)
if is_anomaly and toks > 5: # เริ่มเช็คที่ 5 tokens ขึ้นไป
if attempt < max_retries - 1:
reason = "พูดรัว/อ่านข้าม" if dur_per_tok < min_limit else "เสียงยานคาง/วนลูป"
print(f"⚠️ [{reason}] Detected: {dur_per_tok:.3f}s/tok (Limit: {min_limit:.2f}-{max_limit:.2f}). Retrying ({attempt+1}/3)...")
continue
except Exception as e:
print("Retry check error:", e)
break
yield gr.update(value=output_path, visible=True)
def gen_single_stream(
emo_control_method_value,
prompt,
accent_ref_path,
text,
emo_ref_path,
emo_weight_value,
tvec1_value, tvec2_value, tvec3_value, tvec4_value, tvec5_value,
vec1_value, vec2_value, vec3_value, vec4_value, vec5_value, vec6_value, vec7_value, vec8_value,
emo_text_value,
emo_random_value,
max_text_tokens_per_sentence_value,
duration_seconds_value,
*args,
progress: gr.Progress = gr.Progress()):
if not prompt:
gr.Warning("Upload a prompt audio file first.")
yield None
return
output_path = os.path.join(current_dir, "outputs", f"spk_{int(time.time())}.wav")
try:
tts = ensure_primary_tts()
except RuntimeError as exc:
gr.Warning(str(exc))
yield None
return
tts.gr_progress = progress
advanced_values = list(args)
expected_len = len(advanced_params)
if len(advanced_values) < expected_len:
advanced_values.extend([None] * (expected_len - len(advanced_values)))
raw_generation_kwargs = build_generation_kwargs(*advanced_values[:expected_len])
use_g2p_value = raw_generation_kwargs.pop("use_g2p", False)
use_dataset_spacing_value = raw_generation_kwargs.pop("use_dataset_spacing", False)
trim_silence_value = raw_generation_kwargs.pop("trim_silence", False)
auto_retry_value = raw_generation_kwargs.pop("auto_retry", False)
dur_per_token_value = raw_generation_kwargs.pop("dur_per_token", 0.12)
chain_segments_value = raw_generation_kwargs.pop("chain_segments", False)
generation_kwargs = _prepare_generation_kwargs(raw_generation_kwargs)
generation_kwargs["chain_segments"] = chain_segments_value
emo_mode = emo_control_method_value if isinstance(emo_control_method_value, int) else getattr(emo_control_method_value, "value", 0)
tvec_values = [tvec1_value, tvec2_value, tvec3_value, tvec4_value, tvec5_value]
vec_values = [vec1_value, vec2_value, vec3_value, vec4_value, vec5_value, vec6_value, vec7_value, vec8_value]
if emo_mode == 2:
if sum(tvec_values) > 1.5:
gr.Warning("Thai Emotion vector sum cannot exceed 1.5. Adjust the sliders and retry.")
yield None
return
emo_vector = tvec_values
elif emo_mode == 4:
if sum(vec_values) > 1.5:
gr.Warning("Original Emotion vector sum cannot exceed 1.5. Adjust the sliders and retry.")
yield None
return
emo_vector = vec_values
else:
emo_vector = None
duration_seconds = _normalize_duration_seconds(duration_seconds_value)
preprocessor = ThaiTextPreprocessor(use_g2p=use_g2p_value, use_dataset_spacing=use_dataset_spacing_value)
clean_text = preprocessor.process(text)
if trim_silence_value and prompt and os.path.exists(prompt):
trim_audio_silences(prompt)
# --- Native Gradio Audio Streaming ---
import torchaudio
accumulated_wavs = []
sampling_rate = 22050
ts = int(time.time())
segment_count = 0
chunk_count = 0
# Signal UI to hide normal player, show streaming player, and clear stream buffer
yield None
print("[STREAM DEBUG] Starting infer_generator...")
generator = tts.infer_generator(
spk_audio_prompt=prompt,
text=clean_text,
output_path=None,
emo_audio_prompt=emo_ref_path if emo_mode == 1 else None,
emo_alpha=float(emo_weight_value) if emo_mode in (0, 1) else 1.0,
emo_vector=emo_vector if emo_mode in (2, 4) else None,
use_emo_text=(emo_mode == 3),
emo_text=emo_text_value,
use_random=emo_random_value,
verbose=cmd_args.verbose,
max_text_tokens_per_segment=int(max_text_tokens_per_sentence_value),
duration_seconds=duration_seconds,
stream_return=True,
accent_audio_prompt=accent_ref_path,
**generation_kwargs)
for chunk in generator:
if chunk is None:
print("[STREAM DEBUG] Received None chunk, skipping")
continue
segment_count += 1
accumulated_wavs.append(chunk)
dur_sec = chunk.shape[-1] / sampling_rate
print(f"[STREAM DEBUG] Segment {segment_count}: dur={dur_sec:.2f}s")
if dur_sec < 0.5:
print(f"[STREAM DEBUG] Silence padding, skipping")
continue
chunk_count += 1
audio_np = chunk.squeeze().cpu().numpy()
yield (sampling_rate, audio_np)
print(f"[STREAM DEBUG] Generator done. Segments: {segment_count}, Chunks yielded: {chunk_count}")
# Save final combined file and show in output_audio for replay/download
if accumulated_wavs:
combined = torch.cat(accumulated_wavs, dim=1)
torchaudio.save(output_path, combined.type(torch.int16), sampling_rate)
if trim_silence_value and os.path.exists(output_path):
trim_audio_silences(output_path)
print(f"[STREAM DEBUG] Final audio saved: {output_path}")
# Show final audio + signal JS done
def on_input_text_change(text_value, max_tokens_value):
if not text_value:
return {sentences_preview: gr.update(value=[], visible=True, type="array")}
try:
tts = ensure_primary_tts()
except RuntimeError as exc:
gr.Warning(str(exc))
return {sentences_preview: gr.update(value=[], visible=True, type="array")}
tokenized = tts.tokenizer.tokenize(text_value)
try:
sentences = tts.tokenizer.split_segments(
tokenized, max_text_tokens_per_segment=int(max_tokens_value)
)
data = []
for idx, sentence_tokens in enumerate(sentences):
sentence_str = "".join(sentence_tokens)
data.append([idx, sentence_str, len(sentence_tokens)])
except (AssertionError, Exception) as e:
# Tokenizer assertion: a segment longer than max_text_tokens.
# Show a warning row instead of crashing.
data = [["⚠️", f"Cannot preview: {e}", 0]]
return {sentences_preview: gr.update(value=data, visible=True, type="array")}
def on_method_select(emo_control_value):
if emo_control_value == 0:
return gr.update(visible=False), gr.update(visible=True), gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(visible=False)
if emo_control_value == 1:
return gr.update(visible=True), gr.update(visible=True), gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(visible=False)
if emo_control_value == 2:
return gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(visible=True), gr.update(visible=False), gr.update(visible=False)
if emo_control_value == 3:
return gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(visible=True), gr.update(visible=False)
if emo_control_value == 4:
return gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(visible=True)
return gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(visible=False), gr.update(visible=False)
def build_batch_table_data(rows: List[Dict[str, Any]]):
table_data = []
for row in rows:
text_preview = (row.get("text") or "")[:57]
if row.get("text") and len(row["text"]) > 60:
text_preview += "..."
table_data.append(
[
row.get("id"),
os.path.basename(row.get("prompt_path", "")) if row.get("prompt_path") else "",
text_preview,
os.path.basename(row.get("output_path", "")) if row.get("output_path") else "",
row.get("status", "Pending"),
row.get("last_generated", ""),
]
)
return table_data
def find_batch_row(rows, row_id):
for row in rows or []:
if row.get("id") == row_id:
return row
return None
def resolve_batch_selection(rows, selected_value):
choices = [str(row.get("id")) for row in rows or []]
if not choices:
return gr.update(choices=[], value=None), None
if selected_value is not None:
selected_str = str(selected_value)
if selected_str in choices:
return gr.update(choices=choices, value=selected_str), int(selected_str)
return gr.update(choices=choices, value=choices[-1]), int(choices[-1])
def prepare_batch_selection(rows, selected_value):
dropdown_update, resolved_id = resolve_batch_selection(rows, selected_value)
row = find_batch_row(rows, resolved_id)
prompt_update = gr.update(value=row.get("prompt_path") if row else None)
output_update = gr.update(value=row.get("output_path") if row else None)
text_update = gr.update(value=row.get("text", "") if row else "")
return dropdown_update, resolved_id, prompt_update, output_update, text_update, row
def format_batch_status(row, message=None):
if not row:
base = "No entry selected."
else:
details = [f"Row {row.get('id')}: {row.get('status', 'Pending')}"]
if row.get("text"):
preview = row["text"][:117] + ("..." if len(row["text"]) > 120 else "")
details.append(f"Text: {preview}")
if row.get("output_path"):
details.append(f"Output: {row['output_path']}")
if row.get("last_generated"):
details.append(f"Last generated: {row['last_generated']}")
base = "\n".join(details)
if message:
base = f"{base}\n{message}" if base else message
return gr.update(value=base)
def add_batch_prompts(files, rows, next_id, selected_value):
rows = rows or []
next_id = next_id or 1
files = files or []
updated_rows = [dict(row) for row in rows]
prompts_dir = os.path.join(current_dir, "prompts")
os.makedirs(prompts_dir, exist_ok=True)
added = 0
last_added_id = None
for file_path in files:
if not file_path:
continue
safe_name = os.path.basename(file_path)
timestamp = int(time.time() * 1000)
target_name = f"batch_prompt_{next_id}_{timestamp}_{safe_name}"
target_path = os.path.join(prompts_dir, target_name)
try:
shutil.copy(file_path, target_path)
except Exception as exc:
logger.exception("Failed to store prompt %s", file_path)
gr.Warning(f"Failed to add {safe_name}: {exc}")
continue
entry = {
"id": next_id,
"prompt_path": target_path,
"output_path": None,
"status": "Pending",
"last_generated": "",
"text": "",
}
updated_rows.append(entry)
added += 1
last_added_id = entry["id"]
next_id += 1
selected_seed = last_added_id if added else selected_value
dropdown_update, resolved_id, prompt_update, output_update, text_update, selected_row = prepare_batch_selection(
updated_rows, selected_seed
)
table_update = gr.update(value=build_batch_table_data(updated_rows))
status_message = f"Added {added} prompt{'s' if added != 1 else ''}." if added else "No new prompts were added."
status_update = format_batch_status(selected_row, status_message)
return updated_rows, next_id, gr.update(value=None), table_update, dropdown_update, prompt_update, output_update, text_update, status_update
def validate_emotion_settings(emo_control_method_value, tvec_values, vec_values):
mode = emo_control_method_value if isinstance(emo_control_method_value, int) else getattr(
emo_control_method_value, "value", 0
)
try:
mode = int(mode)
except (TypeError, ValueError):
mode = 0
vec = None
if mode == 2:
if sum(tvec_values) > 1.5:
gr.Warning("Thai vector sum cannot exceed 1.5.")
return mode, None
vec = tvec_values
elif mode == 4:
if sum(vec_values) > 1.5:
gr.Warning("Orig vector sum cannot exceed 1.5.")
return mode, None
vec = vec_values
return mode, vec
def load_dataset_entries(dataset_path, rows, next_id, selected_value, *, progress: Optional[gr.Progress] = None):
rows = rows or []
next_id = next_id or 1
dataset_path = (dataset_path or "").strip()
if not dataset_path:
gr.Warning("Provide a dataset train.txt path before loading.")
dropdown_update, _, prompt_update, output_update, text_update, row = prepare_batch_selection(rows, selected_value)
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(row)
return rows, next_id, gr.update(value=""), table_update, dropdown_update, prompt_update, output_update, text_update, status_update
dataset_path_abs = dataset_path if os.path.isabs(dataset_path) else os.path.abspath(os.path.join(current_dir, dataset_path))
if not os.path.exists(dataset_path_abs):
gr.Warning(f"Dataset file not found: {dataset_path_abs}")
dropdown_update, _, prompt_update, output_update, text_update, row = prepare_batch_selection(rows, selected_value)
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(row)
return rows, next_id, gr.update(value=dataset_path), table_update, dropdown_update, prompt_update, output_update, text_update, status_update
dataset_dir = os.path.dirname(dataset_path_abs)
candidate_dirs = [dataset_dir, os.path.join(dataset_dir, "wavs"), os.path.join(dataset_dir, "audio")]
try:
lines = Path(dataset_path_abs).read_text(encoding="utf-8").splitlines()
except Exception as exc:
gr.Warning(f"Failed to read dataset file: {exc}")
dropdown_update, _, prompt_update, output_update, text_update, row = prepare_batch_selection(rows, selected_value)
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(row)
return rows, next_id, gr.update(value=dataset_path), table_update, dropdown_update, prompt_update, output_update, text_update, status_update
updated_rows = [dict(row) for row in rows]
prompts_dir = os.path.join(current_dir, "prompts")
os.makedirs(prompts_dir, exist_ok=True)
existing_prompts = {os.path.basename(r.get("prompt_path", "")) for r in updated_rows if r.get("prompt_path")}
added = 0
missing_audio = 0
invalid_lines = 0
total_lines = len(lines)
_update_progress(progress, 0.0, desc="Parsing dataset")
for idx, raw_line in enumerate(lines):
_update_progress(progress, min((idx + 1) / max(total_lines, 1), 0.95), desc=f"Processing line {idx + 1}/{total_lines}")
stripped = raw_line.strip()
if not stripped or stripped.startswith("#"):
continue
parts = stripped.split("|", 1)
if len(parts) != 2:
invalid_lines += 1
continue
audio_name = parts[0].strip()
text_value = parts[1].strip()
if not audio_name or not text_value:
invalid_lines += 1
continue
source_path = None
for base_dir in candidate_dirs:
candidate = os.path.join(base_dir, audio_name)
if os.path.exists(candidate):
source_path = candidate
break
if not source_path:
missing_audio += 1
continue
unique_prefix = f"dataset_{next_id}_{int(time.time() * 1000)}"
target_name = f"{unique_prefix}_{os.path.basename(audio_name)}"
if target_name in existing_prompts:
target_name = f"{unique_prefix}_{next_id}_{os.path.basename(audio_name)}"
target_path = os.path.join(prompts_dir, target_name)
try:
shutil.copy(source_path, target_path)
except Exception as exc:
logger.exception("Failed to copy dataset prompt %s", source_path)
gr.Warning(f"Failed to copy {audio_name}: {exc}")
missing_audio += 1
continue
entry = {
"id": next_id,
"prompt_path": target_path,
"output_path": None,
"status": "Pending",
"last_generated": "",
"text": text_value,
}
updated_rows.append(entry)
existing_prompts.add(target_name)
added += 1
next_id += 1
selected_seed = updated_rows[-1]["id"] if added else selected_value
dropdown_update, resolved_id, prompt_update, output_update, text_update, selected_row = prepare_batch_selection(
updated_rows, selected_seed
)
table_update = gr.update(value=build_batch_table_data(updated_rows))
messages = []
if added:
messages.append(f"Loaded {added} entries")
if missing_audio:
messages.append(f"{missing_audio} missing audio")
if invalid_lines:
messages.append(f"{invalid_lines} invalid lines")
status_message = ", ".join(messages) if messages else "No new entries loaded."
status_update = format_batch_status(selected_row, status_message)
_update_progress(progress, 1.0, desc="Dataset load complete")
return updated_rows, next_id, gr.update(value=dataset_path), table_update, dropdown_update, prompt_update, output_update, text_update, status_update
def generate_all_batch(rows, selected_value, worker_count_value, emo_control_method_value, emo_ref_path, emo_weight_value, tvec1_value, tvec2_value, tvec3_value, tvec4_value, tvec5_value, vec1_value, vec2_value, vec3_value, vec4_value, vec5_value, vec6_value, vec7_value, vec8_value, emo_text_value, emo_random_value, max_text_tokens_per_sentence_value, duration_seconds_value, batch_accent_ref, *advanced_param_values, progress: Optional[gr.Progress] = None):
rows = rows or []
if not rows:
gr.Warning("Add prompt audio files before generating.")
dropdown_update, _, prompt_update, output_update, text_update, row = prepare_batch_selection(rows, selected_value)
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(row)
return rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
if parallel_worker_config.get("gpt_path") is None or parallel_worker_config.get("bpe_path") is None:
gr.Warning("Load a GPT checkpoint and BPE tokenizer before generating.")
dropdown_update, _, prompt_update, output_update, text_update, row = prepare_batch_selection(rows, selected_value)
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(row, "Model not loaded.")
return rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
tvec_values = [tvec1_value, tvec2_value, tvec3_value, tvec4_value, tvec5_value]
vec_values = [vec1_value, vec2_value, vec3_value, vec4_value, vec5_value, vec6_value, vec7_value, vec8_value]
emo_mode, emo_vector = validate_emotion_settings(emo_control_method_value, tvec_values, vec_values)
if emo_mode == 2 and emo_vector is None:
dropdown_update, _, prompt_update, output_update, text_update, row = prepare_batch_selection(rows, selected_value)
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(row, "Emotion vector sum exceeded limit.")
return rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
try:
max_tokens = int(max_text_tokens_per_sentence_value)
except (TypeError, ValueError):
max_tokens = 120
duration_seconds = _normalize_duration_seconds(duration_seconds_value)
adv_values = list(advanced_param_values)
expected_len = len(advanced_params)
if len(adv_values) < expected_len:
adv_values.extend([None] * (expected_len - len(adv_values)))
base_generation_kwargs = build_generation_kwargs(*adv_values[:expected_len])
use_g2p_value = base_generation_kwargs.pop("use_g2p", False) # ดึงค่าออกไปใช้
outputs_dir = os.path.join(current_dir, "outputs", "tasks")
os.makedirs(outputs_dir, exist_ok=True)
jobs: List[GenerationJob] = []
row_map: Dict[int, Dict[str, Any]] = {}
for row in rows:
new_row = dict(row)
prompt_path = new_row.get("prompt_path")
if not prompt_path or not os.path.exists(prompt_path):
new_row["status"] = "Error: Prompt missing"
row_map[new_row["id"]] = new_row
continue
text_value = (new_row.get("text") or "").strip()
if not text_value:
new_row["status"] = "Error: Text missing"
row_map[new_row["id"]] = new_row
continue
use_g2p_val = base_generation_kwargs.pop("use_g2p", False)
use_dataset_spacing_val = base_generation_kwargs.pop("use_dataset_spacing", False)
# --- ใช้ Preprocessor ในโหมด Batch ---
preprocessor = ThaiTextPreprocessor(use_g2p=use_g2p_val, use_dataset_spacing=use_dataset_spacing_val)
clean_text = preprocessor.process(text_value)
output_path = os.path.join(outputs_dir, f"batch_row_{new_row['id']}_{int(time.time() * 1000)}.wav")
new_row["status"] = "Running"
new_row["output_path"] = output_path
row_map[new_row["id"]] = new_row
jobs.append(
GenerationJob(
row_id=new_row["id"],
prompt_path=prompt_path,
text=clean_text, # ส่งข้อความที่คลีนแล้วให้ Worker
output_path=output_path,
emo_mode=emo_mode,
emo_weight=float(emo_weight_value) if emo_mode in (0, 1) else 1.0,
emo_vector=emo_vector if emo_mode in (2, 4) else None,
emo_text=emo_text_value,
emo_random=bool(emo_random_value),
emo_ref_path=emo_ref_path if emo_mode == 1 else None,
max_tokens=max_tokens,
generation_kwargs=dict(base_generation_kwargs),
verbose=cmd_args.verbose,
duration_seconds=duration_seconds,
accent_ref_path=batch_accent_ref)
)
running_rows = list(row_map.values())
table_running = gr.update(value=build_batch_table_data(running_rows))
dropdown_update, resolved_id, prompt_update, output_update, text_update, selected_row = prepare_batch_selection(
running_rows, selected_value
)
if not jobs:
status_update = format_batch_status(selected_row, "No rows ready for generation.")
return running_rows, table_running, dropdown_update, prompt_update, output_update, text_update, status_update
_update_progress(progress, 0.0, desc="Starting parallel generation")
worker_pool.ensure(worker_count_value)
results = worker_pool.run_jobs(jobs, progress)
for row_id, result in results.items():
row_entry = row_map.get(row_id)
if not row_entry:
continue
row_entry["status"] = result["status"]
row_entry["last_generated"] = result.get("timestamp", "")
if result["output_path"]:
row_entry["output_path"] = result["output_path"]
final_rows = list(row_map.values())
table_update = gr.update(value=build_batch_table_data(final_rows))
dropdown_update, resolved_id, prompt_update, output_update, text_update, selected_row = prepare_batch_selection(
final_rows, resolved_id
)
status_update = format_batch_status(selected_row, "Parallel generation finished.")
return final_rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
def regenerate_batch_entry(rows, selected_value, worker_count_value, emo_control_method_value, emo_ref_path, emo_weight_value, tvec1_value, tvec2_value, tvec3_value, tvec4_value, tvec5_value, vec1_value, vec2_value, vec3_value, vec4_value, vec5_value, vec6_value, vec7_value, vec8_value, emo_text_value, emo_random_value, max_text_tokens_per_sentence_value, duration_seconds_value, batch_accent_ref, *advanced_param_values, progress: Optional[gr.Progress] = None):
rows = rows or []
dropdown_update, resolved_id, prompt_update, output_update, text_update, selected_row = prepare_batch_selection(rows, selected_value)
if not selected_row:
gr.Warning("Select an entry to regenerate.")
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(None)
return rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
if parallel_worker_config.get("gpt_path") is None or parallel_worker_config.get("bpe_path") is None:
gr.Warning("Load a GPT checkpoint and BPE tokenizer before generating.")
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(selected_row, "Model not loaded.")
return rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
tvec_values = [tvec1_value, tvec2_value, tvec3_value, tvec4_value, tvec5_value]
vec_values = [vec1_value, vec2_value, vec3_value, vec4_value, vec5_value, vec6_value, vec7_value, vec8_value]
emo_mode, emo_vector = validate_emotion_settings(emo_control_method_value, tvec_values, vec_values)
if emo_mode == 2 and emo_vector is None:
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(selected_row, "Emotion vector sum exceeded limit.")
return rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
prompt_path = selected_row.get("prompt_path")
if not prompt_path or not os.path.exists(prompt_path):
gr.Warning("Prompt audio file is missing.")
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(selected_row, "Prompt audio file missing.")
return rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
text_value = (selected_row.get("text") or "").strip()
if not text_value:
gr.Warning("Enter text for this entry before regenerating.")
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(selected_row, "Text is missing.")
return rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
try:
max_tokens = int(max_text_tokens_per_sentence_value)
except (TypeError, ValueError):
max_tokens = 120
adv_values = list(advanced_param_values)
expected_len = len(advanced_params)
if len(adv_values) < expected_len:
adv_values.extend([None] * (expected_len - len(adv_values)))
generation_kwargs = build_generation_kwargs(*adv_values[:expected_len])
use_g2p_value = generation_kwargs.pop("use_g2p", False)
use_dataset_spacing_value = generation_kwargs.pop("use_dataset_spacing", False)
# --- ใช้ Preprocessor ในโหมด Batch ---
preprocessor = ThaiTextPreprocessor(use_g2p=use_g2p_value, use_dataset_spacing=use_dataset_spacing_value)
clean_text = preprocessor.process(text_value)
outputs_dir = os.path.join(current_dir, "outputs", "tasks")
os.makedirs(outputs_dir, exist_ok=True)
output_path = os.path.join(outputs_dir, f"batch_row_{selected_row['id']}_{int(time.time() * 1000)}.wav")
job = GenerationJob(
row_id=selected_row["id"],
prompt_path=prompt_path,
text=clean_text, # ส่งข้อความที่คลีนแล้วให้ Worker
output_path=output_path,
emo_mode=emo_mode,
emo_weight=float(emo_weight_value) if emo_mode in (0, 1) else 1.0,
emo_vector=emo_vector if emo_mode in (2, 4) else None,
emo_text=emo_text_value,
emo_random=bool(emo_random_value),
emo_ref_path=emo_ref_path if emo_mode == 1 else None,
max_tokens=max_tokens,
generation_kwargs=dict(generation_kwargs),
verbose=cmd_args.verbose,
duration_seconds=duration_seconds,
accent_ref_path=batch_accent_ref)
_update_progress(progress, 0.0, desc="Regenerating entry")
worker_pool.ensure(worker_count_value)
results = worker_pool.run_jobs([job], progress)
result = results.get(job.row_id)
updated_rows = []
for row in rows:
if row.get("id") != job.row_id:
updated_rows.append(dict(row))
continue
new_row = dict(row)
if result:
new_row["status"] = result["status"]
new_row["output_path"] = result.get("output_path", new_row.get("output_path"))
new_row["last_generated"] = result.get("timestamp", "")
else:
new_row["status"] = "Error: Unknown"
updated_rows.append(new_row)
table_update = gr.update(value=build_batch_table_data(updated_rows))
dropdown_update, resolved_id, prompt_update, output_update, text_update, selected_row = prepare_batch_selection(
updated_rows, job.row_id
)
status_update = format_batch_status(selected_row, "Regeneration finished.")
return updated_rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
def delete_batch_entry(rows, selected_value):
rows = rows or []
dropdown_update, resolved_id, prompt_update, output_update, text_update, selected_row = prepare_batch_selection(rows, selected_value)
if not selected_row:
gr.Warning("Select an entry to delete.")
table_update = gr.update(value=build_batch_table_data(rows))
status_update = format_batch_status(None)
return rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
remaining_rows = [dict(row) for row in rows if row.get("id") != selected_row.get("id")]
dropdown_update, resolved_id, prompt_update, output_update, text_update, row = prepare_batch_selection(remaining_rows, None)
table_update = gr.update(value=build_batch_table_data(remaining_rows))
status_update = format_batch_status(row, "Entry deleted.")
return remaining_rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
def clear_batch_rows(rows, next_id):
dropdown_update = gr.update(choices=[], value=None)
prompt_update = gr.update(value=None)
output_update = gr.update(value=None)
text_update = gr.update(value="")
status_update = format_batch_status(None, "Batch list cleared.")
return [], 1, gr.update(value=[]), dropdown_update, prompt_update, output_update, text_update, status_update
def on_select_batch_entry(selected_value, rows):
dropdown_update, resolved_id, prompt_update, output_update, text_update, row = prepare_batch_selection(rows, selected_value)
status_update = format_batch_status(row)
return dropdown_update, prompt_update, output_update, text_update, status_update
def update_batch_text(new_text, rows, selected_value):
rows = rows or []
try:
selected_id = int(selected_value) if selected_value is not None else None
except (TypeError, ValueError):
selected_id = None
if selected_id is None:
gr.Warning("Select an entry before editing text.")
table_update = gr.update(value=build_batch_table_data(rows))
dropdown_update, resolved_id, prompt_update, output_update, text_update, row = prepare_batch_selection(rows, selected_value)
status_update = format_batch_status(row)
return rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
updated_rows = []
target_row = None
for row in rows:
new_row = dict(row)
if row.get("id") == selected_id:
new_row["text"] = new_text
if new_row.get("output_path"):
new_row["status"] = "Pending"
target_row = new_row
updated_rows.append(new_row)
dropdown_update, resolved_id, prompt_update, output_update, text_update, row = prepare_batch_selection(updated_rows, selected_id)
table_update = gr.update(value=build_batch_table_data(updated_rows))
status_update = format_batch_status(row, "Text updated. Regenerate to apply." if target_row else None)
return updated_rows, table_update, dropdown_update, prompt_update, output_update, text_update, status_update
def update_prompt_audio():
return gr.update(interactive=True)
emo_control_method.select(
on_method_select,
inputs=[emo_control_method],
outputs=[emotion_reference_group, emo_weight_group, emo_random, thai_emotion_vector_group, emo_text_group, emotion_vector_group])
input_text_single.change(
on_input_text_change,
inputs=[input_text_single, max_text_tokens_per_sentence],
outputs=[sentences_preview])
max_text_tokens_per_sentence.change(
on_input_text_change,
inputs=[input_text_single, max_text_tokens_per_sentence],
outputs=[sentences_preview])
def format_text_action(text_val, use_g2p_val):
if not text_val:
return text_val
preprocessor = ThaiTextPreprocessor(use_g2p=use_g2p_val, use_dataset_spacing=True)
return preprocessor.process(text_val)
format_single_btn.click(
format_text_action,
inputs=[input_text_single, use_g2p],
outputs=[input_text_single])
format_batch_btn.click(
format_text_action,
inputs=[batch_text_input, use_g2p],
outputs=[batch_text_input])
prompt_audio.upload(update_prompt_audio, inputs=[], outputs=[gen_button])
gen_button.click(
gen_single,
inputs=[
emo_control_method,
prompt_audio,
accent_audio,
input_text_single,
emo_upload,
emo_weight,
tvec1,
tvec2,
tvec3,
tvec4,
tvec5,
vec1,
vec2,
vec3,
vec4,
vec5,
vec6,
vec7,
vec8,
emo_text,
emo_random,
max_text_tokens_per_sentence,
duration_seconds_input,
*advanced_params,
],
outputs=[output_audio],
show_progress=True)
gen_stream_button.click(
gen_single_stream,
inputs=[
emo_control_method,
prompt_audio,
accent_audio,
input_text_single,
emo_upload,
emo_weight,
tvec1, tvec2, tvec3, tvec4, tvec5,
vec1, vec2, vec3, vec4, vec5, vec6, vec7, vec8,
emo_text,
emo_random,
max_text_tokens_per_sentence,
duration_seconds_input,
*advanced_params,
],
outputs=[stream_audio_output],
show_progress=True)
# Segment Builder Handlers
def on_select_segment(evt: gr.SelectData, wavs):
idx = evt.index[0] # row index
if idx < len(wavs):
wav_tensor = wavs[idx]
wav_data = wav_tensor.type(torch.int16).numpy().T
return gr.update(value=(22050, wav_data))
return gr.update(value=None)
seg_table.select(on_select_segment, inputs=[seg_state_wavs], outputs=[seg_playback])
def seg_split(text_val, max_tokens, use_g2p_val, use_dataset_spacing_val):
if not text_val:
return [], 0, gr.update(value=[]), gr.update(value="Please enter text.", interactive=False), gr.update(interactive=False)
try:
tts = ensure_primary_tts()
except RuntimeError as exc:
return [], 0, gr.update(value=[]), gr.update(value=str(exc), interactive=False), gr.update(interactive=False)
preprocessor = ThaiTextPreprocessor(use_g2p=use_g2p_val, use_dataset_spacing=use_dataset_spacing_val)
clean_text = preprocessor.process(text_val)
tokenized = tts.tokenizer.tokenize(clean_text)
sentences_tokens = tts.tokenizer.split_segments(tokenized, max_text_tokens_per_segment=int(max_tokens))
sentences = ["".join(s) for s in sentences_tokens]
table_data = [[i+1, s, "Pending", 0.0] for i, s in enumerate(sentences)]
status = f"แบ่งข้อความได้ {len(sentences)} ท่อน พร้อมสำหรับ Generate!"
return sentences, [], 0, gr.update(value=table_data), gr.update(value=status), gr.update(interactive=True), gr.update(interactive=False)
seg_split_btn.click(
seg_split,
inputs=[seg_input_text, max_text_tokens_per_sentence, use_g2p, use_dataset_spacing],
outputs=[seg_state_texts, seg_state_wavs, seg_state_idx, seg_table, seg_status, seg_gen_next_btn, seg_regen_last_btn]
)
seg_format_btn.click(
format_text_action,
inputs=[seg_input_text, use_g2p],
outputs=[seg_input_text])
def seg_generate_chunk(
texts, wavs, idx,
emo_control_method_value,
prompt,
accent_ref_path,
emo_ref_path,
emo_weight_value,
tvec1_value, tvec2_value, tvec3_value, tvec4_value, tvec5_value,
vec1_value, vec2_value, vec3_value, vec4_value, vec5_value, vec6_value, vec7_value, vec8_value,
emo_text_value,
emo_random_value,
max_text_tokens_per_sentence_value,
duration_seconds_value,
*args
):
if not prompt:
return wavs, idx, gr.update(), gr.update(value="Upload a prompt audio file first."), gr.update(), gr.update(), gr.update(), gr.update()
if idx >= len(texts):
return wavs, idx, gr.update(), gr.update(value="สร้างครบทุกท่อนแล้ว! 🎉"), gr.update(), gr.update(), gr.update(), gr.update()
try:
tts = ensure_primary_tts()
except RuntimeError as exc:
return wavs, idx, gr.update(), gr.update(value=str(exc)), gr.update(), gr.update(), gr.update(), gr.update()
advanced_values = list(args)
expected_len = len(advanced_params)
if len(advanced_values) < expected_len:
advanced_values.extend([None] * (expected_len - len(advanced_values)))
raw_generation_kwargs = build_generation_kwargs(*advanced_values[:expected_len])
use_g2p_value = raw_generation_kwargs.pop("use_g2p", False)
use_dataset_spacing_value = raw_generation_kwargs.pop("use_dataset_spacing", False)
trim_silence_value = raw_generation_kwargs.pop("trim_silence", False)
auto_retry_value = raw_generation_kwargs.pop("auto_retry", False)
dur_per_token_value = raw_generation_kwargs.pop("dur_per_token", 0.12)
chain_segments_value = raw_generation_kwargs.pop("chain_segments", False)
generation_kwargs = _prepare_generation_kwargs(raw_generation_kwargs)
generation_kwargs["chain_segments"] = chain_segments_value
emo_mode = emo_control_method_value if isinstance(emo_control_method_value, int) else getattr(emo_control_method_value, "value", 0)
tvec_values = [tvec1_value, tvec2_value, tvec3_value, tvec4_value, tvec5_value]
vec_values = [vec1_value, vec2_value, vec3_value, vec4_value, vec5_value, vec6_value, vec7_value, vec8_value]
if emo_mode == 2:
emo_vector = tvec_values
elif emo_mode == 4:
emo_vector = vec_values
else:
emo_vector = None
# If chain_segments is on and we have previous wavs, we need to pass the last wav as prompt
# But the 'infer' function does this internally if we pass the whole text.
# Since we are passing segment by segment, we must manually handle chaining.
current_prompt = prompt
if chain_segments_value and len(wavs) > 0:
# Save last wav temporarily to use as prompt
import torchaudio
temp_prompt = os.path.join(current_dir, "outputs", "temp_chain_prompt.wav")
last_wav = wavs[-1]
torchaudio.save(temp_prompt, last_wav.type(torch.int16), 22050)
current_prompt = temp_prompt
text = texts[idx]
output_path = os.path.join(current_dir, "outputs", f"seg_{idx}_{int(time.time())}.wav")
# Use infer to get the chunk
tts.infer(
spk_audio_prompt=current_prompt,
text=text,
output_path=output_path,
emo_audio_prompt=emo_ref_path if emo_mode == 1 else None,
emo_alpha=float(emo_weight_value) if emo_mode in (0, 1) else 1.0,
emo_vector=emo_vector if emo_mode in (2, 4) else None,
use_emo_text=(emo_mode == 3),
emo_text=emo_text_value,
use_random=emo_random_value,
verbose=cmd_args.verbose,
max_text_tokens_per_segment=int(max_text_tokens_per_sentence_value),
duration_seconds=_normalize_duration_seconds(duration_seconds_value),
accent_audio_prompt=accent_ref_path,
**generation_kwargs)
import librosa
if trim_silence_value and os.path.exists(output_path):
trim_audio_silences(output_path)
y, sr = librosa.load(output_path, sr=22050)
wav_tensor = torch.tensor(y).unsqueeze(0)
new_wavs = list(wavs)
# Check if this is a regeneration
if idx < len(new_wavs):
new_wavs[idx] = wav_tensor
else:
new_wavs.append(wav_tensor)
# Combine all for final output
combined_tensor = torch.cat(new_wavs, dim=1) if len(new_wavs) > 1 else new_wavs[0]
final_output = os.path.join(current_dir, "outputs", f"combined_{int(time.time())}.wav")
import torchaudio
torchaudio.save(final_output, combined_tensor.type(torch.int16), 22050)
# Update Table
table_data = []
for i, s in enumerate(texts):
status = "Pending"
dur = 0.0
if i < len(new_wavs):
status = "Done"
dur = round(new_wavs[i].shape[1] / 22050, 2)
table_data.append([i+1, s, status, dur])
new_idx = len(new_wavs)
status_msg = f"สร้างท่อนที่ {new_idx} เสร็จแล้ว (จากทั้งหมด {len(texts)} ท่อน)"
has_next = new_idx < len(texts)
has_prev = new_idx > 0
return (
new_wavs,
new_idx,
gr.update(value=table_data),
gr.update(value=status_msg),
gr.update(value=output_path),
gr.update(value=final_output),
gr.update(interactive=has_next),
gr.update(interactive=has_prev)
)
def seg_generate_next(*args):
return seg_generate_chunk(*args)
def seg_regenerate_last(texts, wavs, idx, *args):
# Regenerate the last segment by passing idx - 1
if idx > 0:
return seg_generate_chunk(texts, wavs, idx - 1, *args)
return wavs, idx, gr.update(), gr.update(value="ไม่มีท่อนให้ Regenerate"), gr.update(), gr.update(), gr.update(), gr.update()
def seg_clear():
return [], [], 0, gr.update(value=[]), gr.update(value="ล้างข้อมูลแล้ว"), gr.update(value=None), gr.update(value=None), gr.update(interactive=False), gr.update(interactive=False)
seg_gen_next_btn.click(
seg_generate_next,
inputs=[
seg_state_texts, seg_state_wavs, seg_state_idx,
emo_control_method, seg_prompt_audio, seg_accent_audio, emo_upload, emo_weight,
tvec1, tvec2, tvec3, tvec4, tvec5,
vec1, vec2, vec3, vec4, vec5, vec6, vec7, vec8,
emo_text, emo_random, max_text_tokens_per_sentence, duration_seconds_input,
*advanced_params,
],
outputs=[seg_state_wavs, seg_state_idx, seg_table, seg_status, current_seg_audio, final_seg_audio, seg_gen_next_btn, seg_regen_last_btn]
)
seg_regen_last_btn.click(
seg_regenerate_last,
inputs=[
seg_state_texts, seg_state_wavs, seg_state_idx,
emo_control_method, seg_prompt_audio, seg_accent_audio, emo_upload, emo_weight,
tvec1, tvec2, tvec3, tvec4, tvec5,
vec1, vec2, vec3, vec4, vec5, vec6, vec7, vec8,
emo_text, emo_random, max_text_tokens_per_sentence, duration_seconds_input,
*advanced_params,
],
outputs=[seg_state_wavs, seg_state_idx, seg_table, seg_status, current_seg_audio, final_seg_audio, seg_gen_next_btn, seg_regen_last_btn]
)
seg_clear_btn.click(
seg_clear,
inputs=[],
outputs=[seg_state_texts, seg_state_wavs, seg_state_idx, seg_table, seg_status, current_seg_audio, final_seg_audio, seg_gen_next_btn, seg_regen_last_btn]
)
batch_file_input.upload(
add_batch_prompts,
inputs=[batch_file_input, batch_rows_state, next_batch_id_state, selected_entry],
outputs=[batch_rows_state, next_batch_id_state, batch_file_input, batch_table, selected_entry, batch_prompt_player, batch_output_player, batch_text_input, batch_status])
load_dataset_button.click(
load_dataset_entries,
inputs=[dataset_path_input, batch_rows_state, next_batch_id_state, selected_entry],
outputs=[batch_rows_state, next_batch_id_state, dataset_path_input, batch_table, selected_entry, batch_prompt_player, batch_output_player, batch_text_input, batch_status])
selected_entry.change(
on_select_batch_entry,
inputs=[selected_entry, batch_rows_state],
outputs=[selected_entry, batch_prompt_player, batch_output_player, batch_text_input, batch_status])
apply_text_button.click(
update_batch_text,
inputs=[batch_text_input, batch_rows_state, selected_entry],
outputs=[batch_rows_state, batch_table, selected_entry, batch_prompt_player, batch_output_player, batch_text_input, batch_status])
generate_all_button.click(
generate_all_batch,
inputs=[
batch_rows_state,
selected_entry,
worker_count,
emo_control_method,
emo_upload,
emo_weight,
tvec1,
tvec2,
tvec3,
tvec4,
tvec5,
vec1,
vec2,
vec3,
vec4,
vec5,
vec6,
vec7,
vec8,
emo_text,
emo_random,
max_text_tokens_per_sentence,
duration_seconds_input,
batch_accent_input,
*advanced_params,
],
outputs=[batch_rows_state, batch_table, selected_entry, batch_prompt_player, batch_output_player, batch_text_input, batch_status])
regenerate_button.click(
regenerate_batch_entry,
inputs=[
batch_rows_state,
selected_entry,
worker_count,
emo_control_method,
emo_upload,
emo_weight,
tvec1,
tvec2,
tvec3,
tvec4,
tvec5,
vec1,
vec2,
vec3,
vec4,
vec5,
vec6,
vec7,
vec8,
emo_text,
emo_random,
max_text_tokens_per_sentence,
duration_seconds_input,
batch_accent_input,
*advanced_params,
],
outputs=[batch_rows_state, batch_table, selected_entry, batch_prompt_player, batch_output_player, batch_text_input, batch_status])
delete_entry_button.click(
delete_batch_entry,
inputs=[batch_rows_state, selected_entry],
outputs=[batch_rows_state, batch_table, selected_entry, batch_prompt_player, batch_output_player, batch_text_input, batch_status])
clear_entries_button.click(
clear_batch_rows,
inputs=[batch_rows_state, next_batch_id_state],
outputs=[batch_rows_state, next_batch_id_state, batch_table, selected_entry, batch_prompt_player, batch_output_player, batch_text_input, batch_status])
return demo
def main():
target_gpt = r"C:\datasetmaker\index-tts\models\thaiseperate2.pth"
target_bpe = r"C:\datasetmaker\index-tts\checkpoints\thai_segmented_bpe.model"
if os.path.exists(target_gpt) and os.path.exists(target_bpe):
print(">> Auto-loading default models before UI launch... Please wait.")
try:
load_primary_tts(target_gpt, target_bpe)
print(">> Models auto-loaded successfully!")
except Exception as e:
print(">> Failed to auto-load default models:", e)
demo = create_demo()
demo.queue(20)
print(">> Launching WebUI on http://127.0.0.1:7862")
demo.launch(inbrowser=True, server_name="127.0.0.1", server_port=cmd_args.port,
allowed_paths=[os.path.join(current_dir, "outputs")])
if __name__ == "__main__":
mp.set_start_method("spawn", force=True)
main()