vaghawan commited on
Commit
5827d9c
·
verified ·
1 Parent(s): 0bbcb14

Add multi-speaker 3-speaker epoch-10 inference bundle

Browse files
.gitattributes CHANGED
@@ -33,3 +33,6 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ references/hausa_fe_naijavoices_O0456.wav filter=lfs diff=lfs merge=lfs -text
37
+ references/hausa_fe_waxal_nlp_3.wav filter=lfs diff=lfs merge=lfs -text
38
+ references/hausa_fe_waxal_nlp_5.wav filter=lfs diff=lfs merge=lfs -text
README.md ADDED
@@ -0,0 +1,60 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: other
3
+ language:
4
+ - ha
5
+ tags:
6
+ - text-to-speech
7
+ - xtts
8
+ - hausa
9
+ - coqui
10
+ library_name: coqui-tts
11
+ ---
12
+
13
+ # XTTS-v2 Hausa — multi-speaker (3 speakers, 10 epochs)
14
+
15
+ Fine-tuned [Coqui XTTS-v2](https://huggingface.co/coqui/XTTS-v2) for Hausa (`ha`), trained on
16
+ `vaghawan/hausa-tts-24khz-waxalnlp-3-clean`, `vaghawan/hausa-tts-24khz-waxalnlp-5-clean`, `vaghawan/hausa-tts-24khz-naijavoices-O0456`.
17
+
18
+ ## Files
19
+
20
+ | File | Purpose |
21
+ |------|---------|
22
+ | `best_model.pth` | Fine-tuned GPT checkpoint (epoch 10) |
23
+ | `config.json` | Model config (includes Hausa language) |
24
+ | `vocab.json` | Extended Hausa BPE vocabulary |
25
+ | `references/hausa_fe_waxal_nlp_3.wav` | Speaker reference for `hausa_fe_waxal_nlp_3` |
26
+ | `references/hausa_fe_waxal_nlp_5.wav` | Speaker reference for `hausa_fe_waxal_nlp_5` |
27
+ | `references/hausa_fe_naijavoices_O0456.wav` | Speaker reference for `hausa_fe_naijavoices_O0456` |
28
+ | `infer.py` | Inference script |
29
+ | `xtts_hausa_patch.py` | Required Hausa runtime patches |
30
+ | `env_config.py` | Config loader |
31
+ | `config.env.example` | Example settings (copy to `config.env`) |
32
+ | `samples.txt` | Default eval sentences |
33
+ | `requirements.txt` | Python dependencies |
34
+
35
+ ## Quick start
36
+
37
+ ```bash
38
+ pip install -r requirements.txt
39
+ cp config.env.example config.env
40
+
41
+ python infer.py --model-dir . --speaker-wav references/hausa_fe_waxal_nlp_3.wav
42
+
43
+ # Single line:
44
+ python infer.py --model-dir . --text "Ina kwana." --speaker-wav references/hausa_fe_waxal_nlp_3.wav --out outputs/out.wav
45
+ ```
46
+
47
+ # All speakers (one WAV per reference):
48
+ python infer.py --model-dir . --all-speakers --out-dir outputs/samples
49
+
50
+ ## Training details
51
+
52
+ - Base model: Coqui XTTS-v2
53
+ - Epochs: 10
54
+ - Language: Hausa (`ha`)
55
+ - Speakers: `hausa_fe_waxal_nlp_3`, `hausa_fe_waxal_nlp_5`, `hausa_fe_naijavoices_O0456`
56
+
57
+ ## License
58
+
59
+ XTTS-v2 uses the [Coqui Public Model License](https://huggingface.co/coqui/XTTS-v2).
60
+ Check each dataset license before redistribution.
best_model.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:115a61d67ec6a577ee21eaf3764bdc3304b4959c92c0f062ee7692d793199a92
3
+ size 5646218714
config.env.example ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Inference config for this Hugging Face repo (copy to config.env)
2
+ LANGUAGE=ha
3
+ SPEAKER_IDS=hausa_fe_waxal_nlp_3 hausa_fe_waxal_nlp_5 hausa_fe_naijavoices_O0456
4
+ SPEAKER_REF=references/hausa_fe_waxal_nlp_3.wav
5
+ SAMPLES_FILE=samples.txt
6
+ INFER_OUT_DIR=outputs/samples
7
+ INFER_OUT=outputs/out.wav
8
+ INFER_ALL_SPEAKERS=true
9
+ TEMPERATURE=0.7
10
+ LENGTH_PENALTY=1.0
11
+ REPETITION_PENALTY=5.0
12
+ TOP_K=50
13
+ TOP_P=0.85
config.json ADDED
@@ -0,0 +1,207 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "output_path": "/workspace/mtn-tts-tests/xtts-v2/checkpoints/run/training",
3
+ "logger_uri": null,
4
+ "run_name": "GPT_XTTS_v2_Hausa_FT_multi_speaker",
5
+ "project_name": "XTTS_trainer",
6
+ "run_description": "XTTS-v2 GPT fine-tune on Hausa (WaxalNLP)",
7
+ "print_step": 25,
8
+ "plot_step": 100,
9
+ "model_param_stats": false,
10
+ "wandb_entity": null,
11
+ "dashboard_logger": "tensorboard",
12
+ "save_on_interrupt": true,
13
+ "log_model_step": 100,
14
+ "save_step": 500,
15
+ "save_n_checkpoints": 2,
16
+ "save_checkpoints": true,
17
+ "save_all_best": false,
18
+ "save_best_after": 0,
19
+ "target_loss": null,
20
+ "print_eval": true,
21
+ "test_delay_epochs": 0,
22
+ "run_eval": true,
23
+ "run_eval_steps": null,
24
+ "distributed_backend": "nccl",
25
+ "distributed_url": "tcp://localhost:54321",
26
+ "mixed_precision": false,
27
+ "precision": "fp16",
28
+ "epochs": 10,
29
+ "batch_size": 2,
30
+ "eval_batch_size": 2,
31
+ "grad_clip": 0.0,
32
+ "scheduler_after_epoch": true,
33
+ "lr": 5e-06,
34
+ "optimizer": "AdamW",
35
+ "optimizer_params": {
36
+ "betas": [
37
+ 0.9,
38
+ 0.96
39
+ ],
40
+ "eps": 1e-08,
41
+ "weight_decay": 0.01
42
+ },
43
+ "lr_scheduler": "MultiStepLR",
44
+ "lr_scheduler_params": {
45
+ "milestones": [
46
+ 900000,
47
+ 2700000,
48
+ 5400000
49
+ ],
50
+ "gamma": 0.5,
51
+ "last_epoch": -1
52
+ },
53
+ "use_grad_scaler": false,
54
+ "allow_tf32": false,
55
+ "cudnn_enable": true,
56
+ "cudnn_deterministic": false,
57
+ "cudnn_benchmark": false,
58
+ "training_seed": 1,
59
+ "model": "xtts",
60
+ "num_loader_workers": 4,
61
+ "num_eval_loader_workers": 0,
62
+ "use_noise_augment": false,
63
+ "audio": {
64
+ "sample_rate": 22050,
65
+ "output_sample_rate": 24000,
66
+ "dvae_sample_rate": 22050
67
+ },
68
+ "model_args": {
69
+ "gpt_batch_size": 1,
70
+ "enable_redaction": false,
71
+ "kv_cache": true,
72
+ "gpt_checkpoint": "",
73
+ "clvp_checkpoint": null,
74
+ "decoder_checkpoint": null,
75
+ "num_chars": 255,
76
+ "tokenizer_file": "/workspace/mtn-tts-tests/xtts-v2/checkpoints/XTTS_v2.0_original_model_files/vocab.json",
77
+ "gpt_max_audio_tokens": 605,
78
+ "gpt_max_text_tokens": 402,
79
+ "gpt_max_prompt_tokens": 70,
80
+ "gpt_layers": 30,
81
+ "gpt_n_model_channels": 1024,
82
+ "gpt_n_heads": 16,
83
+ "gpt_number_text_tokens": 8238,
84
+ "gpt_start_text_token": 261,
85
+ "gpt_stop_text_token": 0,
86
+ "gpt_num_audio_tokens": 1026,
87
+ "gpt_start_audio_token": 1024,
88
+ "gpt_stop_audio_token": 1025,
89
+ "gpt_code_stride_len": 1024,
90
+ "gpt_use_masking_gt_prompt_approach": true,
91
+ "gpt_use_perceiver_resampler": true,
92
+ "input_sample_rate": 22050,
93
+ "output_sample_rate": 24000,
94
+ "output_hop_length": 256,
95
+ "decoder_input_dim": 1024,
96
+ "d_vector_dim": 512,
97
+ "cond_d_vector_in_each_upsampling_layer": true,
98
+ "duration_const": 102400,
99
+ "min_conditioning_length": 11025,
100
+ "max_conditioning_length": 132300,
101
+ "gpt_loss_text_ce_weight": 0.01,
102
+ "gpt_loss_mel_ce_weight": 1.0,
103
+ "debug_loading_failures": true,
104
+ "max_wav_length": 330750,
105
+ "max_text_length": 250,
106
+ "mel_norm_file": "/workspace/mtn-tts-tests/xtts-v2/checkpoints/XTTS_v2.0_original_model_files/mel_stats.pth",
107
+ "dvae_checkpoint": "/workspace/mtn-tts-tests/xtts-v2/checkpoints/XTTS_v2.0_original_model_files/dvae.pth",
108
+ "xtts_checkpoint": "/workspace/mtn-tts-tests/xtts-v2/checkpoints/XTTS_v2.0_original_model_files/model.pth",
109
+ "vocoder": ""
110
+ },
111
+ "_supports_cloning": true,
112
+ "use_phonemes": false,
113
+ "phonemizer": null,
114
+ "phoneme_language": null,
115
+ "compute_input_seq_cache": false,
116
+ "text_cleaner": null,
117
+ "enable_eos_bos_chars": false,
118
+ "test_sentences_file": "",
119
+ "phoneme_cache_path": null,
120
+ "characters": null,
121
+ "add_blank": false,
122
+ "batch_group_size": 48,
123
+ "loss_masking": null,
124
+ "min_audio_len": 1,
125
+ "max_audio_len": Infinity,
126
+ "min_text_len": 1,
127
+ "max_text_len": Infinity,
128
+ "compute_f0": false,
129
+ "compute_energy": false,
130
+ "compute_linear_spec": false,
131
+ "precompute_num_workers": 0,
132
+ "start_by_longest": false,
133
+ "shuffle": false,
134
+ "drop_last": false,
135
+ "datasets": [
136
+ {
137
+ "formatter": "",
138
+ "dataset_name": "",
139
+ "path": "",
140
+ "meta_file_train": "",
141
+ "ignored_speakers": null,
142
+ "language": "",
143
+ "phonemizer": "",
144
+ "meta_file_val": "",
145
+ "meta_file_attn_mask": ""
146
+ }
147
+ ],
148
+ "test_sentences": [
149
+ {
150
+ "text": "Ina kwana. Me kuke so mu yi yau?",
151
+ "speaker_wav": [
152
+ "/workspace/mtn-tts-tests/xtts-v2/dataset/references/hausa_fe_waxal_nlp_3.wav"
153
+ ],
154
+ "language": "ha"
155
+ },
156
+ {
157
+ "text": "Za ku iya taimaka mini da saita aikace aikacen katin sim?",
158
+ "speaker_wav": [
159
+ "/workspace/mtn-tts-tests/xtts-v2/dataset/references/hausa_fe_waxal_nlp_3.wav"
160
+ ],
161
+ "language": "ha"
162
+ }
163
+ ],
164
+ "eval_split_max_size": 256,
165
+ "eval_split_size": 0.01,
166
+ "use_speaker_weighted_sampler": false,
167
+ "speaker_weighted_sampler_alpha": 1.0,
168
+ "use_language_weighted_sampler": false,
169
+ "language_weighted_sampler_alpha": 1.0,
170
+ "use_length_weighted_sampler": false,
171
+ "length_weighted_sampler_alpha": 1.0,
172
+ "model_dir": null,
173
+ "languages": [
174
+ "en",
175
+ "es",
176
+ "fr",
177
+ "de",
178
+ "it",
179
+ "pt",
180
+ "pl",
181
+ "tr",
182
+ "ru",
183
+ "nl",
184
+ "cs",
185
+ "ar",
186
+ "zh-cn",
187
+ "hu",
188
+ "ko",
189
+ "ja",
190
+ "hi",
191
+ "ha"
192
+ ],
193
+ "temperature": 0.85,
194
+ "length_penalty": 1.0,
195
+ "repetition_penalty": 2.0,
196
+ "top_k": 50,
197
+ "top_p": 0.85,
198
+ "num_gpt_outputs": 1,
199
+ "gpt_cond_len": 12,
200
+ "gpt_cond_chunk_len": 4,
201
+ "max_ref_len": 10,
202
+ "sound_norm_refs": false,
203
+ "optimizer_wd_only_on_weights": true,
204
+ "weighted_loss_attrs": {},
205
+ "weighted_loss_multipliers": {},
206
+ "github_branch": "* main"
207
+ }
env_config.py ADDED
@@ -0,0 +1,182 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Load KEY=VALUE settings from config.env (project root).
3
+
4
+ CLI flags always win over config.env values.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import os
10
+ import shlex
11
+ from pathlib import Path
12
+
13
+ ROOT = Path(__file__).resolve().parent
14
+ DEFAULT_ENV_PATH = ROOT / "config.env"
15
+
16
+ # Always take these from config.env so a broken shell mirror/token cannot win.
17
+ _CONFIG_FORCE_KEYS = frozenset(
18
+ {
19
+ "HF_ENDPOINT",
20
+ "HF_TOKEN",
21
+ "HUGGINGFACE_ACCESS_TOKEN",
22
+ "HUGGING_FACE_HUB_TOKEN",
23
+ "HUGGINGFACE_TOKEN",
24
+ "HF_HUB_DOWNLOAD_TIMEOUT",
25
+ "HF_HUB_ETAG_TIMEOUT",
26
+ "HF_DOWNLOAD_RETRIES",
27
+ }
28
+ )
29
+
30
+
31
+ def load_config_env(path: Path | None = None, *, override: bool = False) -> Path | None:
32
+ """Parse config.env into os.environ.
33
+
34
+ Returns the path loaded, or None if missing.
35
+ Does not override existing environment variables unless override=True,
36
+ except for `_CONFIG_FORCE_KEYS` (HF endpoint/token/timeouts).
37
+ """
38
+ env_path = (path or DEFAULT_ENV_PATH).resolve()
39
+ if not env_path.is_file():
40
+ return None
41
+
42
+ for raw in env_path.read_text(encoding="utf-8").splitlines():
43
+ line = raw.strip()
44
+ if not line or line.startswith("#"):
45
+ continue
46
+ if line.startswith("export "):
47
+ line = line[len("export ") :].strip()
48
+ if "=" not in line:
49
+ continue
50
+ key, value = line.split("=", 1)
51
+ key = key.strip()
52
+ value = value.strip()
53
+ if not key:
54
+ continue
55
+ if len(value) >= 2 and value[0] == value[-1] and value[0] in "\"'":
56
+ value = value[1:-1]
57
+ # Skip blank assignments (e.g. FOO=) so they don't poison Hub clients
58
+ if value == "":
59
+ continue
60
+ if override or key not in os.environ or key in _CONFIG_FORCE_KEYS:
61
+ os.environ[key] = value
62
+ return env_path
63
+
64
+
65
+ def env_str(name: str, default: str | None = None) -> str | None:
66
+ v = os.getenv(name)
67
+ if v is None or v == "":
68
+ return default
69
+ return v
70
+
71
+
72
+ def env_int(name: str, default: int) -> int:
73
+ v = os.getenv(name)
74
+ if v is None or v == "":
75
+ return default
76
+ return int(v)
77
+
78
+
79
+ def env_float(name: str, default: float) -> float:
80
+ v = os.getenv(name)
81
+ if v is None or v == "":
82
+ return default
83
+ return float(v)
84
+
85
+
86
+ def env_path(name: str, default: str | Path) -> Path:
87
+ v = os.getenv(name)
88
+ return Path(v) if v else Path(default)
89
+
90
+
91
+ def env_list(name: str, default: list[str] | None = None) -> list[str]:
92
+ """Split a config value into a list (whitespace or comma separated)."""
93
+ v = os.getenv(name)
94
+ if v is None or v.strip() == "":
95
+ return list(default or [])
96
+ normalized = v.replace(",", " ")
97
+ parts = shlex.split(normalized)
98
+ return [p for p in parts if p]
99
+
100
+
101
+ def ensure_config_loaded(path: Path | None = None) -> None:
102
+ loaded = load_config_env(path)
103
+ if loaded:
104
+ print(f"[config] loaded {loaded}")
105
+ else:
106
+ example = ROOT / "config.env.example"
107
+ print(
108
+ f"[config] no config.env found"
109
+ + (f" (copy from {example.name})" if example.is_file() else "")
110
+ )
111
+ apply_hf_token()
112
+
113
+
114
+ def _sync_hf_hub_endpoint(endpoint: str) -> None:
115
+ """Keep huggingface_hub / datasets in sync if they were already imported."""
116
+ endpoint = endpoint.rstrip("/")
117
+ try:
118
+ import huggingface_hub.constants as hub_constants
119
+
120
+ hub_constants.ENDPOINT = endpoint
121
+ hub_constants.HUGGINGFACE_CO_URL_TEMPLATE = endpoint + "/{repo_id}/resolve/{revision}/{filename}"
122
+ except Exception:
123
+ pass
124
+ try:
125
+ import datasets.config as ds_config
126
+
127
+ ds_config.HF_ENDPOINT = endpoint
128
+ ds_config.HUB_DATASETS_URL = endpoint + "/datasets/{repo_id}/resolve/{revision}/{path}"
129
+ except Exception:
130
+ pass
131
+
132
+
133
+ def apply_hf_token() -> None:
134
+ """Normalize HF token aliases and sync Hub endpoint.
135
+
136
+ HF_ENDPOINT from config.env already wins over the shell via
137
+ `_CONFIG_FORCE_KEYS` — no need to `unset HF_ENDPOINT` before runs.
138
+ """
139
+ endpoint = (os.environ.get("HF_ENDPOINT") or "").strip().rstrip("/")
140
+ if endpoint and not endpoint.startswith(("http://", "https://")):
141
+ del os.environ["HF_ENDPOINT"]
142
+ endpoint = ""
143
+ if not endpoint:
144
+ endpoint = "https://huggingface.co"
145
+ os.environ["HF_ENDPOINT"] = endpoint
146
+
147
+ _sync_hf_hub_endpoint(endpoint)
148
+ print(f"[config] HF_ENDPOINT={endpoint}")
149
+
150
+ # Longer timeouts help IncompleteRead on slow/unstable links
151
+ for key, default in (
152
+ ("HF_HUB_DOWNLOAD_TIMEOUT", "120"),
153
+ ("HF_HUB_ETAG_TIMEOUT", "60"),
154
+ ("HF_DOWNLOAD_RETRIES", "8"),
155
+ ):
156
+ val = env_str(key, default)
157
+ if val:
158
+ os.environ[key] = val
159
+
160
+ token = (
161
+ env_str("HF_TOKEN")
162
+ or env_str("HUGGINGFACE_ACCESS_TOKEN")
163
+ or env_str("HUGGING_FACE_HUB_TOKEN")
164
+ or env_str("HUGGINGFACE_TOKEN")
165
+ )
166
+ if not token:
167
+ print("[config] HF token not set (optional for public repos)")
168
+ return
169
+
170
+ os.environ["HF_TOKEN"] = token
171
+ os.environ["HUGGING_FACE_HUB_TOKEN"] = token
172
+ if not os.getenv("HUGGINGFACE_ACCESS_TOKEN"):
173
+ os.environ["HUGGINGFACE_ACCESS_TOKEN"] = token
174
+
175
+ print(f"[config] HF token loaded ({token[:6]}…{token[-4:] if len(token) > 10 else '****'})")
176
+
177
+ try:
178
+ from huggingface_hub import login
179
+
180
+ login(token=token, add_to_git_credential=False)
181
+ except Exception as e:
182
+ print(f"[config] huggingface_hub.login skipped: {e}")
infer.py ADDED
@@ -0,0 +1,279 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Inference with a fine-tuned XTTS-v2 checkpoint.
3
+
4
+ Defaults (from config.env):
5
+ python infer.py
6
+ -> synthesizes every line in SAMPLES_FILE for SPEAKER_REF
7
+
8
+ Examples:
9
+ python infer.py
10
+ python infer.py --samples samples.txt --all-speakers
11
+ python infer.py --text "Ina kwana." --speaker-wav dataset/references/hausa_fe_waxal_nlp_3.wav
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import argparse
17
+ import re
18
+ from pathlib import Path
19
+
20
+ import torch
21
+ import torchaudio
22
+ from TTS.tts.configs.xtts_config import XttsConfig
23
+ from TTS.tts.models.xtts import Xtts
24
+
25
+ from env_config import (
26
+ ensure_config_loaded,
27
+ env_float,
28
+ env_int,
29
+ env_list,
30
+ env_path,
31
+ env_str,
32
+ )
33
+ from xtts_hausa_patch import apply_xtts_hausa_patches
34
+
35
+
36
+ def _env_bool(name: str, default: bool = False) -> bool:
37
+ v = env_str(name)
38
+ if v is None or v == "":
39
+ return default
40
+ return v.lower() in {"1", "true", "yes", "y", "on"}
41
+
42
+
43
+ def _find_latest_run(training_root: Path) -> Path:
44
+ runs = sorted(
45
+ [p for p in training_root.glob("*") if p.is_dir()],
46
+ key=lambda p: p.stat().st_mtime,
47
+ reverse=True,
48
+ )
49
+ if not runs:
50
+ raise SystemExit(f"No training runs found under {training_root}")
51
+ return runs[0]
52
+
53
+
54
+ def _load_samples(path: Path) -> list[str]:
55
+ lines = []
56
+ for raw in path.read_text(encoding="utf-8").splitlines():
57
+ text = raw.strip()
58
+ if text and not text.startswith("#"):
59
+ lines.append(text)
60
+ if not lines:
61
+ raise SystemExit(f"No texts found in {path}")
62
+ return lines
63
+
64
+
65
+ def _safe_stem(text: str, idx: int) -> str:
66
+ slug = re.sub(r"[^a-zA-Z0-9]+", "_", text.lower()).strip("_")
67
+ slug = (slug[:40] or "utt").rstrip("_")
68
+ return f"{idx:02d}_{slug}"
69
+
70
+
71
+ def _resolve_speaker_refs(
72
+ speaker_wav: Path | None,
73
+ all_speakers: bool,
74
+ dataset_dir: Path,
75
+ ) -> list[tuple[str, Path]]:
76
+ """Return list of (speaker_label, wav_path)."""
77
+ if all_speakers:
78
+ ids = env_list("SPEAKER_IDS")
79
+ refs_dir = dataset_dir / "references"
80
+ pairs: list[tuple[str, Path]] = []
81
+ for spk in ids:
82
+ p = refs_dir / f"{spk}.wav"
83
+ if not p.is_file():
84
+ raise SystemExit(f"Missing reference for speaker {spk}: {p}")
85
+ pairs.append((spk, p))
86
+ if not pairs:
87
+ raise SystemExit("INFER_ALL_SPEAKERS/SPEAKER_IDS set but no references found")
88
+ return pairs
89
+
90
+ if speaker_wav is None:
91
+ raise SystemExit("Provide --speaker-wav or set SPEAKER_REF in config.env")
92
+ if not speaker_wav.is_file():
93
+ raise SystemExit(f"Speaker reference not found: {speaker_wav}")
94
+ label = speaker_wav.stem
95
+ return [(label, speaker_wav)]
96
+
97
+
98
+ def _load_model(model_dir: Path, base_dir: Path) -> tuple[Xtts, Path]:
99
+ config_path = model_dir / "config.json"
100
+ if not config_path.is_file():
101
+ candidates = list(model_dir.rglob("config.json"))
102
+ if not candidates:
103
+ raise SystemExit(f"No config.json under {model_dir}")
104
+ config_path = candidates[0]
105
+ model_dir = config_path.parent
106
+
107
+ ckpt = None
108
+ for name in ("best_model.pth", "model.pth"):
109
+ p = model_dir / name
110
+ if p.is_file():
111
+ ckpt = p
112
+ break
113
+ if ckpt is None:
114
+ pths = sorted(model_dir.glob("*.pth"), key=lambda p: p.stat().st_mtime, reverse=True)
115
+ if not pths:
116
+ raise SystemExit(f"No .pth checkpoint in {model_dir}")
117
+ ckpt = pths[0]
118
+
119
+ vocab = model_dir / "vocab.json"
120
+ if not vocab.is_file():
121
+ vocab = base_dir / "vocab.json"
122
+
123
+ print(f"Loading config={config_path}")
124
+ print(f"Loading checkpoint={ckpt}")
125
+ print(f"Using vocab={vocab}")
126
+
127
+ config = XttsConfig()
128
+ config.load_json(str(config_path))
129
+ model = Xtts.init_from_config(config)
130
+ model.load_checkpoint(
131
+ config,
132
+ checkpoint_path=str(ckpt),
133
+ vocab_path=str(vocab),
134
+ eval=True,
135
+ use_deepspeed=False,
136
+ )
137
+ if torch.cuda.is_available():
138
+ model.cuda()
139
+ return model, model_dir
140
+
141
+
142
+ def _synthesize(
143
+ model: Xtts,
144
+ text: str,
145
+ language: str,
146
+ speaker_wav: Path,
147
+ args: argparse.Namespace,
148
+ ) -> torch.Tensor:
149
+ gpt_cond_latent, speaker_embedding = model.get_conditioning_latents(
150
+ audio_path=str(speaker_wav.resolve()),
151
+ gpt_cond_len=model.config.gpt_cond_len,
152
+ max_ref_length=model.config.max_ref_len,
153
+ sound_norm_refs=model.config.sound_norm_refs,
154
+ )
155
+ out = model.inference(
156
+ text=text,
157
+ language=language,
158
+ gpt_cond_latent=gpt_cond_latent,
159
+ speaker_embedding=speaker_embedding,
160
+ temperature=args.temperature,
161
+ length_penalty=args.length_penalty,
162
+ repetition_penalty=args.repetition_penalty,
163
+ top_k=args.top_k,
164
+ top_p=args.top_p,
165
+ )
166
+ return torch.tensor(out["wav"]).unsqueeze(0)
167
+
168
+
169
+ def main() -> None:
170
+ ensure_config_loaded()
171
+ apply_xtts_hausa_patches()
172
+
173
+ speaker_default = env_str("SPEAKER_REF")
174
+ text_default = env_str("INFER_TEXT")
175
+ samples_default = env_str("SAMPLES_FILE", "samples.txt")
176
+
177
+ ap = argparse.ArgumentParser(description="XTTS-v2 Hausa inference")
178
+ ap.add_argument(
179
+ "--text",
180
+ default=None,
181
+ help="Single utterance (overrides samples file)",
182
+ )
183
+ ap.add_argument(
184
+ "--samples",
185
+ type=Path,
186
+ default=None,
187
+ help="Text file, one utterance per line (default: SAMPLES_FILE from config.env)",
188
+ )
189
+ ap.add_argument(
190
+ "--speaker-wav",
191
+ type=Path,
192
+ default=Path(speaker_default) if speaker_default else None,
193
+ help="Reference WAV (or set SPEAKER_REF in config.env)",
194
+ )
195
+ ap.add_argument(
196
+ "--all-speakers",
197
+ action="store_true",
198
+ default=_env_bool("INFER_ALL_SPEAKERS", False),
199
+ help="Synthesize for every SPEAKER_IDS reference under dataset/references/",
200
+ )
201
+ ap.add_argument("--language", default=env_str("LANGUAGE", "ha"))
202
+ ap.add_argument("--model-dir", type=Path, default=None)
203
+ ap.add_argument(
204
+ "--base-model-dir",
205
+ type=Path,
206
+ default=env_path("BASE_MODEL_DIR", "checkpoints/XTTS_v2.0_original_model_files"),
207
+ )
208
+ ap.add_argument(
209
+ "--out",
210
+ type=Path,
211
+ default=env_path("INFER_OUT", "outputs/out.wav"),
212
+ help="Output path for single --text mode",
213
+ )
214
+ ap.add_argument(
215
+ "--out-dir",
216
+ type=Path,
217
+ default=env_path("INFER_OUT_DIR", "outputs/samples"),
218
+ help="Output directory for samples-file mode",
219
+ )
220
+ ap.add_argument("--temperature", type=float, default=env_float("TEMPERATURE", 0.7))
221
+ ap.add_argument("--length-penalty", type=float, default=env_float("LENGTH_PENALTY", 1.0))
222
+ ap.add_argument("--repetition-penalty", type=float, default=env_float("REPETITION_PENALTY", 5.0))
223
+ ap.add_argument("--top-k", type=int, default=env_int("TOP_K", 50))
224
+ ap.add_argument("--top-p", type=float, default=env_float("TOP_P", 0.85))
225
+ args = ap.parse_args()
226
+
227
+ # Resolve texts: explicit --text > INFER_TEXT (only if --samples not passed) > samples file
228
+ texts: list[str]
229
+ batch_mode: bool
230
+ if args.text:
231
+ texts = [args.text]
232
+ batch_mode = False
233
+ elif args.samples is not None:
234
+ texts = _load_samples(args.samples)
235
+ batch_mode = True
236
+ elif text_default:
237
+ texts = [text_default]
238
+ batch_mode = False
239
+ else:
240
+ samples_path = Path(samples_default) if samples_default else Path("samples.txt")
241
+ texts = _load_samples(samples_path)
242
+ batch_mode = True
243
+ print(f"[infer] using samples file: {samples_path.resolve()}")
244
+
245
+ dataset_dir = env_path("DATASET_DIR", "dataset")
246
+ speakers = _resolve_speaker_refs(args.speaker_wav, args.all_speakers, dataset_dir)
247
+
248
+ base_dir = args.base_model_dir.resolve()
249
+ if args.model_dir is None:
250
+ model_dir = _find_latest_run(Path("checkpoints/run/training").resolve())
251
+ else:
252
+ model_dir = args.model_dir.resolve()
253
+
254
+ model, _ = _load_model(model_dir, base_dir)
255
+
256
+ if batch_mode:
257
+ out_dir = args.out_dir.resolve()
258
+ out_dir.mkdir(parents=True, exist_ok=True)
259
+ for spk_label, spk_wav in speakers:
260
+ spk_dir = out_dir / spk_label if len(speakers) > 1 else out_dir
261
+ spk_dir.mkdir(parents=True, exist_ok=True)
262
+ print(f"\n=== speaker={spk_label} ref={spk_wav} ===")
263
+ for i, text in enumerate(texts, start=1):
264
+ out_path = spk_dir / f"{_safe_stem(text, i)}.wav"
265
+ print(f"[{i}/{len(texts)}] {text[:80]}...")
266
+ wav = _synthesize(model, text, args.language, spk_wav, args)
267
+ torchaudio.save(str(out_path), wav, 24000)
268
+ print(f" -> {out_path}")
269
+ print(f"\nDone. Wrote {len(texts) * len(speakers)} files under {out_dir}")
270
+ else:
271
+ spk_label, spk_wav = speakers[0]
272
+ args.out.parent.mkdir(parents=True, exist_ok=True)
273
+ wav = _synthesize(model, texts[0], args.language, spk_wav, args)
274
+ torchaudio.save(str(args.out), wav, 24000)
275
+ print(f"Wrote {args.out} (speaker={spk_label})")
276
+
277
+
278
+ if __name__ == "__main__":
279
+ main()
references/hausa_fe_naijavoices_O0456.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d722651a022295d6cd4313982249c12c1a57f8ca046d96a71f0a40933fac651e
3
+ size 377326
references/hausa_fe_waxal_nlp_3.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e898d3f4fe314fb41e98fa708dfa5e53a348588885bf5c77dc6f18127669e2a0
3
+ size 383854
references/hausa_fe_waxal_nlp_5.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6baa1fa3e4b09898e187bf0089075d62633efab998f7d92a06a09f686959a12c
3
+ size 384044
requirements.txt ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Maintained Coqui TTS fork (includes XTTS GPT fine-tuning trainer)
2
+ coqui-tts[all,codec]>=0.27.0
3
+ torch
4
+ torchaudio
5
+ torchcodec
6
+ transformers>=4.40,<5.0
7
+ datasets>=2.19.0
8
+ soundfile
9
+ librosa
10
+ numpy
11
+ pandas
12
+ tokenizers
13
+ huggingface_hub
14
+ tqdm
15
+ tensorboard
samples.txt ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ za ku iya taimaka min da sa ita aikace aikacen katin sim?
2
+ na san game da katin sim. na iya taimaka maka samun bayyane game da yadda za a iya amfani da shi.menene kuma kake buƙata game da katin sim?
3
+ to, na sanar da kai cewa ba zan iya taimaka maka ba, amma zan iya bayar da shawara, ko taimakon da zai taimake ka. menene kuma bayanan da kake nema?
4
+ kin kaɗ ɗina ba ya kunna ko da bayan saka menene matsalar?
5
+ shin akwai hanyar yin canjin sim na wani daga wayata?
vocab.json ADDED
The diff for this file is too large to render. See raw diff
 
xtts_hausa_patch.py ADDED
@@ -0,0 +1,160 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Runtime patches required for Hausa (and other non-native XTTS languages).
3
+
4
+ 1) VoiceBpeTokenizer.preprocess_text raises NotImplementedError for `ha`.
5
+ 2) After vocab extension, text embeddings grow; load_checkpoint must pad
6
+ pretrained weights instead of skipping the whole embedding matrix.
7
+ 3) Xtts.synthesize asserts language ∈ config.languages — GPTTrainerConfig
8
+ defaults omit langs added only in on-disk config.json (e.g. ha).
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import logging
14
+
15
+ import torch
16
+ from TTS.tts.layers.xtts.tokenizer import VoiceBpeTokenizer
17
+ from TTS.tts.layers.xtts.trainer.gpt_trainer import GPTTrainer
18
+ from TTS.tts.models.xtts import Xtts
19
+ from TTS.tts.utils.text.cleaners import basic_cleaners, collapse_whitespace, lowercase
20
+
21
+ logger = logging.getLogger("xtts_hausa_patch")
22
+ _APPLIED = False
23
+
24
+
25
+ def _hausa_preprocess_text(self, txt, lang):
26
+ lang = (lang or "").split("-")[0]
27
+ supported = {"ar", "cs", "de", "en", "es", "fr", "hi", "hu", "it", "nl", "pl", "pt", "ru", "tr", "zh", "ko"}
28
+ if lang in supported:
29
+ from TTS.tts.layers.xtts.tokenizer import multilingual_cleaners, chinese_transliterate
30
+
31
+ txt = multilingual_cleaners(txt, lang)
32
+ if lang == "zh":
33
+ txt = chinese_transliterate(txt)
34
+ if lang == "ko":
35
+ from ko_speech_tools import hangul_romanize
36
+
37
+ txt = hangul_romanize(txt)
38
+ return txt
39
+ if lang == "ja":
40
+ from TTS.tts.layers.xtts.tokenizer import japanese_cleaners
41
+
42
+ return japanese_cleaners(txt, self.katsu)
43
+
44
+ # New / unsupported languages (Hausa, etc.): keep Latin orthography, light clean.
45
+ txt = txt.replace('"', "")
46
+ txt = lowercase(txt)
47
+ txt = basic_cleaners(txt)
48
+ txt = collapse_whitespace(txt)
49
+ return txt
50
+
51
+
52
+ def _pad_text_layers(state: dict, model) -> dict:
53
+ """Pad gpt text_embedding / text_head when vocab was extended."""
54
+ emb_key = None
55
+ for key in ("gpt.text_embedding.weight", "text_embedding.weight"):
56
+ if key in state:
57
+ emb_key = key
58
+ break
59
+ if emb_key is None or model.gpt is None:
60
+ return state
61
+
62
+ ckpt_emb = state[emb_key]
63
+ model_emb = model.gpt.text_embedding.weight
64
+ if ckpt_emb.shape == model_emb.shape:
65
+ return state
66
+
67
+ if ckpt_emb.shape[0] > model_emb.shape[0]:
68
+ logger.warning(
69
+ "Checkpoint text vocab (%d) larger than model (%d); truncating.",
70
+ ckpt_emb.shape[0],
71
+ model_emb.shape[0],
72
+ )
73
+ state[emb_key] = ckpt_emb[: model_emb.shape[0]].clone()
74
+ return state
75
+
76
+ num_new = model_emb.shape[0] - ckpt_emb.shape[0]
77
+ logger.info("Padding text embedding with %d new token rows for extended vocab.", num_new)
78
+
79
+ new_rows = torch.randn(num_new, ckpt_emb.shape[1], dtype=ckpt_emb.dtype)
80
+ # Keep old EOS/special last-row convention used by Coqui GPTTrainer.
81
+ start_token_row = ckpt_emb[-1, :].clone()
82
+ padded = torch.cat([ckpt_emb, new_rows], dim=0)
83
+ padded[-1, :] = start_token_row
84
+ state[emb_key] = padded
85
+
86
+ head_w_key = "gpt.text_head.weight" if "gpt.text_head.weight" in state else "text_head.weight"
87
+ head_b_key = "gpt.text_head.bias" if "gpt.text_head.bias" in state else "text_head.bias"
88
+ if head_w_key in state:
89
+ tw = state[head_w_key]
90
+ start = tw[-1, :].clone()
91
+ new_w = torch.randn(num_new, tw.shape[1], dtype=tw.dtype)
92
+ tw = torch.cat([tw, new_w], dim=0)
93
+ tw[-1, :] = start
94
+ state[head_w_key] = tw
95
+ if head_b_key in state:
96
+ tb = state[head_b_key]
97
+ start_b = tb[-1].clone()
98
+ new_b = torch.zeros(num_new, dtype=tb.dtype)
99
+ tb = torch.cat([tb, new_b], dim=0)
100
+ tb[-1] = start_b
101
+ state[head_b_key] = tb
102
+
103
+ return state
104
+
105
+
106
+ _ORIG_LOAD_CHECKPOINT = GPTTrainer.load_checkpoint
107
+ _ORIG_SYNTHESIZE = Xtts.synthesize
108
+
109
+
110
+ def _ensure_language(config, language: str | None) -> None:
111
+ """Allow fine-tune languages that exist in vocab but not in dataclass defaults."""
112
+ if not language or not hasattr(config, "languages"):
113
+ return
114
+ lang = "zh-cn" if language == "zh" else language
115
+ langs = list(config.languages or [])
116
+ if lang not in langs:
117
+ langs.append(lang)
118
+ config.languages = langs
119
+ logger.info("Registered language '%s' on XTTS config for synthesize/inference.", lang)
120
+
121
+
122
+ def _patched_synthesize(self, text, config=None, *, speaker_wav=None, language=None, **kwargs):
123
+ _ensure_language(self.config, language)
124
+ return _ORIG_SYNTHESIZE(
125
+ self, text, config, speaker_wav=speaker_wav, language=language, **kwargs
126
+ )
127
+
128
+
129
+ def _patched_load_checkpoint(
130
+ self,
131
+ config,
132
+ checkpoint_path,
133
+ *,
134
+ eval=False,
135
+ strict=True,
136
+ cache_storage="/tmp/tts_cache",
137
+ target_protocol="s3",
138
+ target_options=None,
139
+ ):
140
+ if target_options is None:
141
+ target_options = {"anon": True}
142
+ state = self.xtts.get_compatible_checkpoint_state_dict(checkpoint_path)
143
+ state = _pad_text_layers(state, self.xtts)
144
+ # After padding, shapes match — prefer strict load of GPT weights when possible.
145
+ self.xtts.load_state_dict(state, strict=False)
146
+ if eval:
147
+ self.xtts.gpt.init_gpt_for_inference(kv_cache=self.args.kv_cache, use_deepspeed=False)
148
+ self.eval()
149
+
150
+
151
+ def apply_xtts_hausa_patches() -> None:
152
+ global _APPLIED
153
+ if _APPLIED:
154
+ return
155
+ VoiceBpeTokenizer.preprocess_text = _hausa_preprocess_text
156
+ GPTTrainer.load_checkpoint = _patched_load_checkpoint
157
+ Xtts.synthesize = _patched_synthesize
158
+ _APPLIED = True
159
+ logger.info("Applied XTTS Hausa patches (tokenizer + embedding pad + languages).")
160
+ print("[patch] XTTS Hausa tokenizer + embedding-pad + language patches applied")