Add multi-speaker 3-speaker epoch-10 inference bundle
Browse files- .gitattributes +3 -0
- README.md +60 -0
- best_model.pth +3 -0
- config.env.example +13 -0
- config.json +207 -0
- env_config.py +182 -0
- infer.py +279 -0
- references/hausa_fe_naijavoices_O0456.wav +3 -0
- references/hausa_fe_waxal_nlp_3.wav +3 -0
- references/hausa_fe_waxal_nlp_5.wav +3 -0
- requirements.txt +15 -0
- samples.txt +5 -0
- vocab.json +0 -0
- xtts_hausa_patch.py +160 -0
.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")
|