#!/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()