File size: 5,714 Bytes
1352e38
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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")