vaghawan's picture
Add English/Pidgin XTTS-v2 FT bundle (lr=5e-5, epoch 20, 35 mins)
379b63f verified
Raw History Blame Contribute Delete
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")