Download infer.py from vaghawan/xtts-v2-en-pidgin-lr-5e-5-epoch10-35mins: direct link, hf CLI and curl.
- Browser
- Download file 9.18 kB
-
https://huggingface.co/vaghawan/xtts-v2-en-pidgin-lr-5e-5-epoch10-35mins/resolve/main/infer.py
- Command line
-
hf download hf://vaghawan/xtts-v2-en-pidgin-lr-5e-5-epoch10-35mins/infer.py
-
curl -L -o infer.py https://huggingface.co/vaghawan/xtts-v2-en-pidgin-lr-5e-5-epoch10-35mins/resolve/main/infer.py
9.18 kB
| #!/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() | |