Download xtts_hausa_patch.py from vaghawan/xtts-v2-en-pidgin-lr-5e-5-epoch20-35mins: direct link, hf CLI and curl.
- Browser
- Download file 5.71 kB
-
https://huggingface.co/vaghawan/xtts-v2-en-pidgin-lr-5e-5-epoch20-35mins/resolve/main/xtts_hausa_patch.py
- Command line
-
hf download hf://vaghawan/xtts-v2-en-pidgin-lr-5e-5-epoch20-35mins/xtts_hausa_patch.py
-
curl -L -o xtts_hausa_patch.py https://huggingface.co/vaghawan/xtts-v2-en-pidgin-lr-5e-5-epoch20-35mins/resolve/main/xtts_hausa_patch.py
5.71 kB
| #!/usr/bin/env python3 | |
| """Runtime patches required for Hausa (and other non-native XTTS languages). | |
| 1) VoiceBpeTokenizer.preprocess_text raises NotImplementedError for `ha`. | |
| 2) After vocab extension, text embeddings grow; load_checkpoint must pad | |
| pretrained weights instead of skipping the whole embedding matrix. | |
| 3) Xtts.synthesize asserts language ∈ config.languages — GPTTrainerConfig | |
| defaults omit langs added only in on-disk config.json (e.g. ha). | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| import torch | |
| from TTS.tts.layers.xtts.tokenizer import VoiceBpeTokenizer | |
| from TTS.tts.layers.xtts.trainer.gpt_trainer import GPTTrainer | |
| from TTS.tts.models.xtts import Xtts | |
| from TTS.tts.utils.text.cleaners import basic_cleaners, collapse_whitespace, lowercase | |
| logger = logging.getLogger("xtts_hausa_patch") | |
| _APPLIED = False | |
| def _hausa_preprocess_text(self, txt, lang): | |
| lang = (lang or "").split("-")[0] | |
| supported = {"ar", "cs", "de", "en", "es", "fr", "hi", "hu", "it", "nl", "pl", "pt", "ru", "tr", "zh", "ko"} | |
| if lang in supported: | |
| from TTS.tts.layers.xtts.tokenizer import multilingual_cleaners, chinese_transliterate | |
| txt = multilingual_cleaners(txt, lang) | |
| if lang == "zh": | |
| txt = chinese_transliterate(txt) | |
| if lang == "ko": | |
| from ko_speech_tools import hangul_romanize | |
| txt = hangul_romanize(txt) | |
| return txt | |
| if lang == "ja": | |
| from TTS.tts.layers.xtts.tokenizer import japanese_cleaners | |
| return japanese_cleaners(txt, self.katsu) | |
| # New / unsupported languages (Hausa, etc.): keep Latin orthography, light clean. | |
| txt = txt.replace('"', "") | |
| txt = lowercase(txt) | |
| txt = basic_cleaners(txt) | |
| txt = collapse_whitespace(txt) | |
| return txt | |
| def _pad_text_layers(state: dict, model) -> dict: | |
| """Pad gpt text_embedding / text_head when vocab was extended.""" | |
| emb_key = None | |
| for key in ("gpt.text_embedding.weight", "text_embedding.weight"): | |
| if key in state: | |
| emb_key = key | |
| break | |
| if emb_key is None or model.gpt is None: | |
| return state | |
| ckpt_emb = state[emb_key] | |
| model_emb = model.gpt.text_embedding.weight | |
| if ckpt_emb.shape == model_emb.shape: | |
| return state | |
| if ckpt_emb.shape[0] > model_emb.shape[0]: | |
| logger.warning( | |
| "Checkpoint text vocab (%d) larger than model (%d); truncating.", | |
| ckpt_emb.shape[0], | |
| model_emb.shape[0], | |
| ) | |
| state[emb_key] = ckpt_emb[: model_emb.shape[0]].clone() | |
| return state | |
| num_new = model_emb.shape[0] - ckpt_emb.shape[0] | |
| logger.info("Padding text embedding with %d new token rows for extended vocab.", num_new) | |
| new_rows = torch.randn(num_new, ckpt_emb.shape[1], dtype=ckpt_emb.dtype) | |
| # Keep old EOS/special last-row convention used by Coqui GPTTrainer. | |
| start_token_row = ckpt_emb[-1, :].clone() | |
| padded = torch.cat([ckpt_emb, new_rows], dim=0) | |
| padded[-1, :] = start_token_row | |
| state[emb_key] = padded | |
| head_w_key = "gpt.text_head.weight" if "gpt.text_head.weight" in state else "text_head.weight" | |
| head_b_key = "gpt.text_head.bias" if "gpt.text_head.bias" in state else "text_head.bias" | |
| if head_w_key in state: | |
| tw = state[head_w_key] | |
| start = tw[-1, :].clone() | |
| new_w = torch.randn(num_new, tw.shape[1], dtype=tw.dtype) | |
| tw = torch.cat([tw, new_w], dim=0) | |
| tw[-1, :] = start | |
| state[head_w_key] = tw | |
| if head_b_key in state: | |
| tb = state[head_b_key] | |
| start_b = tb[-1].clone() | |
| new_b = torch.zeros(num_new, dtype=tb.dtype) | |
| tb = torch.cat([tb, new_b], dim=0) | |
| tb[-1] = start_b | |
| state[head_b_key] = tb | |
| return state | |
| _ORIG_LOAD_CHECKPOINT = GPTTrainer.load_checkpoint | |
| _ORIG_SYNTHESIZE = Xtts.synthesize | |
| def _ensure_language(config, language: str | None) -> None: | |
| """Allow fine-tune languages that exist in vocab but not in dataclass defaults.""" | |
| if not language or not hasattr(config, "languages"): | |
| return | |
| lang = "zh-cn" if language == "zh" else language | |
| langs = list(config.languages or []) | |
| if lang not in langs: | |
| langs.append(lang) | |
| config.languages = langs | |
| logger.info("Registered language '%s' on XTTS config for synthesize/inference.", lang) | |
| def _patched_synthesize(self, text, config=None, *, speaker_wav=None, language=None, **kwargs): | |
| _ensure_language(self.config, language) | |
| return _ORIG_SYNTHESIZE( | |
| self, text, config, speaker_wav=speaker_wav, language=language, **kwargs | |
| ) | |
| def _patched_load_checkpoint( | |
| self, | |
| config, | |
| checkpoint_path, | |
| *, | |
| eval=False, | |
| strict=True, | |
| cache_storage="/tmp/tts_cache", | |
| target_protocol="s3", | |
| target_options=None, | |
| ): | |
| if target_options is None: | |
| target_options = {"anon": True} | |
| state = self.xtts.get_compatible_checkpoint_state_dict(checkpoint_path) | |
| state = _pad_text_layers(state, self.xtts) | |
| # After padding, shapes match — prefer strict load of GPT weights when possible. | |
| self.xtts.load_state_dict(state, strict=False) | |
| if eval: | |
| self.xtts.gpt.init_gpt_for_inference(kv_cache=self.args.kv_cache, use_deepspeed=False) | |
| self.eval() | |
| def apply_xtts_hausa_patches() -> None: | |
| global _APPLIED | |
| if _APPLIED: | |
| return | |
| VoiceBpeTokenizer.preprocess_text = _hausa_preprocess_text | |
| GPTTrainer.load_checkpoint = _patched_load_checkpoint | |
| Xtts.synthesize = _patched_synthesize | |
| _APPLIED = True | |
| logger.info("Applied XTTS Hausa patches (tokenizer + embedding pad + languages).") | |
| print("[patch] XTTS Hausa tokenizer + embedding-pad + language patches applied") | |