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