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