#!/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")