#!/usr/bin/env python3 """Stock XTTS-v2 inference for English + Nigerian Pidgin (same English reference voice). Uses the base model only (no fine-tune) and clones: dataset/references/english_reference.wav (auto-converted from english_reference.m4a if the WAV is missing). Pidgin has no native XTTS language id, so both groups use language=`en`. Examples: python infer_english_and_pidgin_base.py python infer_english_and_pidgin_base.py --out-dir outputs/samples_en_pidgin_base """ from __future__ import annotations import argparse import shutil import subprocess from pathlib import Path import torchaudio from env_config import ensure_config_loaded, env_float, env_int, env_path from infer import _safe_stem, _synthesize from infer_base import _load_stock_xtts from xtts_hausa_patch import apply_xtts_hausa_patches, _ensure_language ROOT = Path(__file__).resolve().parent DEFAULT_REF_M4A = ROOT / "dataset/references/english_reference.m4a" DEFAULT_REF_WAV = ROOT / "dataset/references/english_reference.wav" DEFAULT_OUT_DIR = ROOT / "outputs/samples_en_pidgin_base" # (group, index, text) — index matches the user's sample numbering SAMPLES: list[tuple[str, int, str]] = [ ( "english", 1, "Hello I am Zigi, your customer support agent. How can I help you today?", ), ( "english", 2, "MTN is a telecommunications company, always by your side.", ), ( "english", 3, ( "MTN is a leading telecommunications company. MTN operates across multiple " "regions, and MTN's network coverage is one of MTN's greatest strengths. " "When MTN launched MTN Mobile Money, MTN changed the way millions of people " "transact. MTN's CEO has stated that MTN will continue to expand MTN's " "infrastructure so that MTN can serve more MTN customers." ), ), ( "pidgin", 4, ( "The Inflight Roaming data service no get limit and e cost N5,000, but you " "fit use am only for some selected airlines." ), ), ( "pidgin", 5, ( "Your data don finish. To buy new one, dial *131# or enter any MTN shop wey " "dey near you. You wan make I help you with anything else?" ), ), ( "pidgin", 6, ( "Thank you for call MTN customer care. We sabi say e important to dey " "connected. Your current plan get unlimited calls inside MTN network and " "five gigabytes of data wey go last thirty days. If you wan upgrade your " "plan or you wan hear about our new offers, our team dey available " "twenty-four seven to help you." ), ), ] def _ensure_reference_wav(m4a: Path, wav: Path) -> Path: """Prefer existing WAV; otherwise convert m4a → 24 kHz mono WAV via ffmpeg.""" if wav.is_file() and wav.stat().st_size > 0: return wav if not m4a.is_file(): raise SystemExit(f"Missing English reference: {m4a} (and no {wav})") wav.parent.mkdir(parents=True, exist_ok=True) ffmpeg = shutil.which("ffmpeg") if not ffmpeg: raise SystemExit("ffmpeg not found; convert english_reference.m4a to .wav first") print(f"[ref] converting {m4a.name} -> {wav.name} (24 kHz mono)") subprocess.run( [ffmpeg, "-y", "-i", str(m4a), "-ac", "1", "-ar", "24000", str(wav)], check=True, capture_output=True, ) return wav def main() -> None: ensure_config_loaded() apply_xtts_hausa_patches() ap = argparse.ArgumentParser( description="Stock XTTS-v2: English + Pidgin samples with english_reference voice" ) ap.add_argument( "--speaker-wav", type=Path, default=DEFAULT_REF_WAV, help="English reference WAV (default: dataset/references/english_reference.wav)", ) ap.add_argument( "--speaker-m4a", type=Path, default=DEFAULT_REF_M4A, help="Source m4a used if --speaker-wav is missing", ) 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-dir", type=Path, default=DEFAULT_OUT_DIR, help="Output directory (default: outputs/samples_en_pidgin_base)", ) ap.add_argument( "--language", default="en", help="XTTS language code (default: en; Pidgin also runs as en)", ) 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() speaker_wav = _ensure_reference_wav(args.speaker_m4a.resolve(), args.speaker_wav.resolve()) if not speaker_wav.is_file(): raise SystemExit(f"Speaker reference not found: {speaker_wav}") base_dir = args.base_model_dir.resolve() if not (base_dir / "model.pth").is_file(): raise SystemExit( f"Stock base model missing under {base_dir}\n" "Run: python download_base_model.py" ) model = _load_stock_xtts(base_dir) _ensure_language(model.config, args.language) out_dir = args.out_dir.resolve() out_dir.mkdir(parents=True, exist_ok=True) print(f"[base] speaker_ref={speaker_wav}") print(f"[base] language={args.language} out_dir={out_dir}") for group, idx, text in SAMPLES: group_dir = out_dir / group group_dir.mkdir(parents=True, exist_ok=True) out_path = group_dir / f"{_safe_stem(text, idx)}.wav" print(f"\n=== [{group}] #{idx} ===") print(f" {text[:100]}{'...' if len(text) > 100 else ''}") wav = _synthesize(model, text, args.language, speaker_wav, args) torchaudio.save(str(out_path), wav, 24000) print(f" -> {out_path}") print(f"\nDone. Wrote {len(SAMPLES)} files under {out_dir}") if __name__ == "__main__": main()