File size: 5,714 Bytes
379b63f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 | #!/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")
|