vaghawan's picture
Add English/Pidgin XTTS-v2 FT bundle (lr=5e-5, epoch 20, 35 mins)
379b63f verified
Raw History Blame Contribute Delete
9.18 kB
#!/usr/bin/env python3
"""Inference with a fine-tuned XTTS-v2 checkpoint.
Defaults (from config.env):
python infer.py
-> synthesizes every line in SAMPLES_FILE for SPEAKER_REF
Examples:
python infer.py
python infer.py --samples samples.txt --all-speakers
python infer.py --text "Ina kwana." --speaker-wav dataset/references/hausa_fe_waxal_nlp_3.wav
"""
from __future__ import annotations
import argparse
import re
from pathlib import Path
import torch
import torchaudio
from TTS.tts.configs.xtts_config import XttsConfig
from TTS.tts.models.xtts import Xtts
from env_config import (
ensure_config_loaded,
env_float,
env_int,
env_list,
env_path,
env_str,
)
from xtts_hausa_patch import apply_xtts_hausa_patches
def _env_bool(name: str, default: bool = False) -> bool:
v = env_str(name)
if v is None or v == "":
return default
return v.lower() in {"1", "true", "yes", "y", "on"}
def _find_latest_run(training_root: Path) -> Path:
runs = sorted(
[p for p in training_root.glob("*") if p.is_dir()],
key=lambda p: p.stat().st_mtime,
reverse=True,
)
if not runs:
raise SystemExit(f"No training runs found under {training_root}")
return runs[0]
def _load_samples(path: Path) -> list[str]:
lines = []
for raw in path.read_text(encoding="utf-8").splitlines():
text = raw.strip()
if text and not text.startswith("#"):
lines.append(text)
if not lines:
raise SystemExit(f"No texts found in {path}")
return lines
def _safe_stem(text: str, idx: int) -> str:
slug = re.sub(r"[^a-zA-Z0-9]+", "_", text.lower()).strip("_")
slug = (slug[:40] or "utt").rstrip("_")
return f"{idx:02d}_{slug}"
def _resolve_speaker_refs(
speaker_wav: Path | None,
all_speakers: bool,
dataset_dir: Path,
) -> list[tuple[str, Path]]:
"""Return list of (speaker_label, wav_path)."""
if all_speakers:
ids = env_list("SPEAKER_IDS")
refs_dir = dataset_dir / "references"
pairs: list[tuple[str, Path]] = []
for spk in ids:
p = refs_dir / f"{spk}.wav"
if not p.is_file():
raise SystemExit(f"Missing reference for speaker {spk}: {p}")
pairs.append((spk, p))
if not pairs:
raise SystemExit("INFER_ALL_SPEAKERS/SPEAKER_IDS set but no references found")
return pairs
if speaker_wav is None:
raise SystemExit("Provide --speaker-wav or set SPEAKER_REF in config.env")
if not speaker_wav.is_file():
raise SystemExit(f"Speaker reference not found: {speaker_wav}")
label = speaker_wav.stem
return [(label, speaker_wav)]
def _load_model(model_dir: Path, base_dir: Path) -> tuple[Xtts, Path]:
config_path = model_dir / "config.json"
if not config_path.is_file():
candidates = list(model_dir.rglob("config.json"))
if not candidates:
raise SystemExit(f"No config.json under {model_dir}")
config_path = candidates[0]
model_dir = config_path.parent
ckpt = None
for name in ("best_model.pth", "model.pth"):
p = model_dir / name
if p.is_file():
ckpt = p
break
if ckpt is None:
pths = sorted(model_dir.glob("*.pth"), key=lambda p: p.stat().st_mtime, reverse=True)
if not pths:
raise SystemExit(f"No .pth checkpoint in {model_dir}")
ckpt = pths[0]
vocab = model_dir / "vocab.json"
if not vocab.is_file():
vocab = base_dir / "vocab.json"
print(f"Loading config={config_path}")
print(f"Loading checkpoint={ckpt}")
print(f"Using vocab={vocab}")
config = XttsConfig()
config.load_json(str(config_path))
model = Xtts.init_from_config(config)
model.load_checkpoint(
config,
checkpoint_path=str(ckpt),
vocab_path=str(vocab),
eval=True,
use_deepspeed=False,
)
if torch.cuda.is_available():
model.cuda()
return model, model_dir
def _synthesize(
model: Xtts,
text: str,
language: str,
speaker_wav: Path,
args: argparse.Namespace,
) -> torch.Tensor:
gpt_cond_latent, speaker_embedding = model.get_conditioning_latents(
audio_path=str(speaker_wav.resolve()),
gpt_cond_len=model.config.gpt_cond_len,
max_ref_length=model.config.max_ref_len,
sound_norm_refs=model.config.sound_norm_refs,
)
out = model.inference(
text=text,
language=language,
gpt_cond_latent=gpt_cond_latent,
speaker_embedding=speaker_embedding,
temperature=args.temperature,
length_penalty=args.length_penalty,
repetition_penalty=args.repetition_penalty,
top_k=args.top_k,
top_p=args.top_p,
)
return torch.tensor(out["wav"]).unsqueeze(0)
def main() -> None:
ensure_config_loaded()
apply_xtts_hausa_patches()
speaker_default = env_str("SPEAKER_REF")
text_default = env_str("INFER_TEXT")
samples_default = env_str("SAMPLES_FILE", "samples.txt")
ap = argparse.ArgumentParser(description="XTTS-v2 Hausa inference")
ap.add_argument(
"--text",
default=None,
help="Single utterance (overrides samples file)",
)
ap.add_argument(
"--samples",
type=Path,
default=None,
help="Text file, one utterance per line (default: SAMPLES_FILE from config.env)",
)
ap.add_argument(
"--speaker-wav",
type=Path,
default=Path(speaker_default) if speaker_default else None,
help="Reference WAV (or set SPEAKER_REF in config.env)",
)
ap.add_argument(
"--all-speakers",
action="store_true",
default=_env_bool("INFER_ALL_SPEAKERS", False),
help="Synthesize for every SPEAKER_IDS reference under dataset/references/",
)
ap.add_argument("--language", default=env_str("LANGUAGE", "ha"))
ap.add_argument("--model-dir", type=Path, default=None)
ap.add_argument(
"--base-model-dir",
type=Path,
default=env_path("BASE_MODEL_DIR", "checkpoints/XTTS_v2.0_original_model_files"),
)
ap.add_argument(
"--out",
type=Path,
default=env_path("INFER_OUT", "outputs/out.wav"),
help="Output path for single --text mode",
)
ap.add_argument(
"--out-dir",
type=Path,
default=env_path("INFER_OUT_DIR", "outputs/samples"),
help="Output directory for samples-file mode",
)
ap.add_argument("--temperature", type=float, default=env_float("TEMPERATURE", 0.7))
ap.add_argument("--length-penalty", type=float, default=env_float("LENGTH_PENALTY", 1.0))
ap.add_argument("--repetition-penalty", type=float, default=env_float("REPETITION_PENALTY", 5.0))
ap.add_argument("--top-k", type=int, default=env_int("TOP_K", 50))
ap.add_argument("--top-p", type=float, default=env_float("TOP_P", 0.85))
args = ap.parse_args()
# Resolve texts: explicit --text > INFER_TEXT (only if --samples not passed) > samples file
texts: list[str]
batch_mode: bool
if args.text:
texts = [args.text]
batch_mode = False
elif args.samples is not None:
texts = _load_samples(args.samples)
batch_mode = True
elif text_default:
texts = [text_default]
batch_mode = False
else:
samples_path = Path(samples_default) if samples_default else Path("samples.txt")
texts = _load_samples(samples_path)
batch_mode = True
print(f"[infer] using samples file: {samples_path.resolve()}")
dataset_dir = env_path("DATASET_DIR", "dataset")
speakers = _resolve_speaker_refs(args.speaker_wav, args.all_speakers, dataset_dir)
base_dir = args.base_model_dir.resolve()
if args.model_dir is None:
model_dir = _find_latest_run(Path("checkpoints/run/training").resolve())
else:
model_dir = args.model_dir.resolve()
model, _ = _load_model(model_dir, base_dir)
if batch_mode:
out_dir = args.out_dir.resolve()
out_dir.mkdir(parents=True, exist_ok=True)
for spk_label, spk_wav in speakers:
spk_dir = out_dir / spk_label if len(speakers) > 1 else out_dir
spk_dir.mkdir(parents=True, exist_ok=True)
print(f"\n=== speaker={spk_label} ref={spk_wav} ===")
for i, text in enumerate(texts, start=1):
out_path = spk_dir / f"{_safe_stem(text, i)}.wav"
print(f"[{i}/{len(texts)}] {text[:80]}...")
wav = _synthesize(model, text, args.language, spk_wav, args)
torchaudio.save(str(out_path), wav, 24000)
print(f" -> {out_path}")
print(f"\nDone. Wrote {len(texts) * len(speakers)} files under {out_dir}")
else:
spk_label, spk_wav = speakers[0]
args.out.parent.mkdir(parents=True, exist_ok=True)
wav = _synthesize(model, texts[0], args.language, spk_wav, args)
torchaudio.save(str(args.out), wav, 24000)
print(f"Wrote {args.out} (speaker={spk_label})")
if __name__ == "__main__":
main()