#!/usr/bin/env python3 """Fine-tuned XTTS-v2 inference for English + Nigerian Pidgin. Requires an explicit --model-dir (never picks “latest” Hausa run by mtime). Defaults: - language=en (Pidgin has no XTTS language id) - speaker: SPEAKER_REF from config.en_pidgin.env, else dataset_en_pidgin/references/ - same 6 hardcoded samples as infer_english_and_pidgin_base.py Examples: python infer_english_and_pidgin.py \\ --config-env config.en_pidgin.env \\ --model-dir checkpoints/run/training/GPT_XTTS_v2_EN_Pidgin_FT- python infer_english_and_pidgin.py \\ --model-dir /path/to/run \\ --speaker-wav dataset_en_pidgin/references/en_pidgin_voice.wav \\ --out-dir outputs/samples_en_pidgin """ from __future__ import annotations import argparse from pathlib import Path import torchaudio from env_config import ensure_config_loaded, env_float, env_int, env_path, env_str from infer import _load_model, _safe_stem, _synthesize from infer_english_and_pidgin_base import SAMPLES, _ensure_reference_wav from xtts_hausa_patch import apply_xtts_hausa_patches, _ensure_language ROOT = Path(__file__).resolve().parent DEFAULT_OUT_DIR = ROOT / "outputs/samples_en_pidgin" DEFAULT_REF = ROOT / "dataset_en_pidgin/references/en_pidgin_voice.wav" DEFAULT_REF_M4A = ROOT / "dataset/references/english_reference.m4a" DEFAULT_REF_FALLBACK_WAV = ROOT / "dataset/references/english_reference.wav" def main() -> None: ensure_config_loaded() apply_xtts_hausa_patches() speaker_default = env_str("SPEAKER_REF") ap = argparse.ArgumentParser( description="Fine-tuned XTTS-v2: English + Pidgin samples (explicit --model-dir)" ) ap.add_argument( "--config-env", type=Path, default=None, help="Alternate config (e.g. config.en_pidgin.env). Also: XTTS_CONFIG_ENV.", ) ap.add_argument( "--model-dir", type=Path, required=True, help="Fine-tuned run directory containing best_model.pth / config.json (required)", ) ap.add_argument( "--base-model-dir", type=Path, default=env_path("BASE_MODEL_DIR", "checkpoints/XTTS_v2.0_stock_en"), help="Fallback vocab/config parent if missing from --model-dir", ) ap.add_argument( "--speaker-wav", type=Path, default=Path(speaker_default) if speaker_default else DEFAULT_REF, help="Reference WAV for voice cloning", ) ap.add_argument( "--out-dir", type=Path, default=Path(env_str("INFER_OUT_DIR", str(DEFAULT_OUT_DIR))), help="Output directory (default: outputs/samples_en_pidgin)", ) ap.add_argument( "--language", default=env_str("LANGUAGE", "en"), help="XTTS language code (default: 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() model_dir = args.model_dir.resolve() if not model_dir.is_dir(): raise SystemExit(f"--model-dir is not a directory: {model_dir}") speaker_wav = args.speaker_wav.resolve() if not speaker_wav.is_file(): # Fall back to english_reference (m4a→wav) used for base A/B speaker_wav = _ensure_reference_wav( DEFAULT_REF_M4A.resolve(), DEFAULT_REF_FALLBACK_WAV.resolve(), ) if not speaker_wav.is_file(): raise SystemExit( f"Speaker reference not found: {args.speaker_wav}\n" "Prepare dataset_en_pidgin/ or pass --speaker-wav" ) model, resolved_dir = _load_model(model_dir, args.base_model_dir.resolve()) _ensure_language(model.config, args.language) out_dir = args.out_dir.resolve() out_dir.mkdir(parents=True, exist_ok=True) print(f"[ft] model_dir={resolved_dir}") print(f"[ft] speaker_ref={speaker_wav}") print(f"[ft] 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()