vaghawan's picture
Add English/Pidgin XTTS-v2 FT bundle (lr=5e-5, epoch 10, 35 mins)
1352e38 verified
Raw History Blame Contribute Delete
7.18 kB
#!/usr/bin/env python3
"""Inference with the stock XTTS-v2 base model (no fine-tune).
Use this as the A/B baseline against `infer.py` (fine-tuned) on the same
`samples.txt` + speaker reference WAVs.
Defaults (from config.env):
python infer_base.py
-> synthesizes every line in SAMPLES_FILE for SPEAKER_REF
-> writes under outputs/samples_base/
Examples:
python infer_base.py
python infer_base.py --samples samples.txt --all-speakers
python infer_base.py --text "Ina kwana." --speaker-wav dataset/references/hausa_fe_waxal_nlp_3.wav
Notes:
- Loads original `model.pth` plus pristine `config.json.bak` / `vocab.json.bak`
when present (pre-`extend_vocab.py`), so weights match stock XTTS-v2.
- Still applies Hausa runtime patches so `language=ha` can run on stock XTTS.
"""
from __future__ import annotations
import argparse
import shutil
import tempfile
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_path,
env_str,
)
from infer import (
_env_bool,
_load_samples,
_resolve_speaker_refs,
_safe_stem,
_synthesize,
)
from xtts_hausa_patch import apply_xtts_hausa_patches, _ensure_language
def _pick_stock_file(base_dir: Path, name: str) -> Path:
"""Prefer pristine *.bak from before extend_vocab; else the live file."""
bak = base_dir / f"{name}.bak"
live = base_dir / name
if bak.is_file():
return bak
if live.is_file():
return live
raise SystemExit(f"Missing {name} (and {name}.bak) under {base_dir}")
def _load_stock_xtts(base_dir: Path) -> Xtts:
"""Load stock XTTS-v2 weights (not a fine-tuned run)."""
base_dir = base_dir.resolve()
ckpt = base_dir / "model.pth"
if not ckpt.is_file():
raise SystemExit(f"Stock checkpoint not found: {ckpt}")
config_src = _pick_stock_file(base_dir, "config.json")
vocab_src = _pick_stock_file(base_dir, "vocab.json")
print(f"[base] Loading stock XTTS-v2 from {base_dir}")
print(f"[base] config={config_src.name} vocab={vocab_src.name} ckpt={ckpt.name}")
if config_src.suffix == ".bak" or vocab_src.suffix == ".bak":
print("[base] Using pre-extend_vocab backups (true stock tokenizer/config)")
# Xtts.load_checkpoint resolves vocab next to checkpoint_dir; stage bak files
# into a temp dir so we never point at the extended live vocab by accident.
with tempfile.TemporaryDirectory(prefix="xtts_stock_") as tmp:
tmp_dir = Path(tmp)
config_path = tmp_dir / "config.json"
vocab_path = tmp_dir / "vocab.json"
shutil.copy2(config_src, config_path)
shutil.copy2(vocab_src, vocab_path)
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_path),
eval=True,
use_deepspeed=False,
)
if torch.cuda.is_available():
model.cuda()
return model
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="Stock XTTS-v2 inference (no fine-tune) for A/B vs infer.py"
)
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)",
)
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(
"--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=Path("outputs/out_base.wav"),
help="Output path for single --text mode",
)
ap.add_argument(
"--out-dir",
type=Path,
default=Path("outputs/samples_base"),
help="Output directory for samples-file mode (default: outputs/samples_base)",
)
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()
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"[base] 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)
model = _load_stock_xtts(args.base_model_dir.resolve())
_ensure_language(model.config, args.language)
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=== [stock] 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)} stock 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} (stock XTTS, speaker={spk_label})")
if __name__ == "__main__":
main()