tktts_qwen3.py ダウンロード/コピー

tktts_qwen3.py をダウンロード

tktts_qwen3.py
tktts_qwen3.py
  1"""tkttsライブラリファミリー向けのQwen3-TTSバックエンド。
  2
  3このモジュールは、公式の ``qwen-tts`` Pythonパッケージを介してQwen3-TTSをローカルで実行しながら、
  4``tktts_voicevox.py`` と互換性のあるインターフェースを提供します。
  5モデルは遅延ロードされキャッシュされるため、連続した呼び出しで再ロードされることはありません。
  6
  7関連リンク:
  8    :doc:`tktts_base`
  9"""
 10
 11from __future__ import annotations
 12
 13import os
 14import re
 15from pathlib import Path
 16from typing import Any
 17
 18
 19missing = []
 20for lib, import_name in [
 21    ("torch", "torch"),
 22    ("soundfile", "soundfile"),
 23    ("qwen-tts", "qwen_tts"),
 24]:
 25    try:
 26        __import__(import_name)
 27    except ImportError:
 28        missing.append(lib)
 29
 30if missing:
 31    raise ImportError(
 32        "Error: Missing libraries: "
 33        + ", ".join(missing)
 34        + "\n  install: pip install qwen-tts soundfile"
 35    )
 36
 37import soundfile as sf
 38import torch
 39from qwen_tts import Qwen3TTSModel
 40
 41try:
 42    from tktts_base import apply_replacements, normalize_speaker, split_dialogue
 43except ImportError as exc:
 44    raise ImportError(
 45        "tktts_qwen3 requires tktts_base.py to be importable."
 46    ) from exc
 47
 48
 49TTS_ENGINE_NAME = "qwen3"
 50DEFAULT_QWEN3_MODEL = "Qwen/Qwen3-TTS-12Hz-0.6B-CustomVoice"
 51DEFAULT_QWEN3_VOICE = "Ono_Anna"
 52DEFAULT_LANGUAGE = "Japanese"
 53DEFAULT_DEVICE = "auto"
 54DEFAULT_DTYPE = "auto"
 55
 56
 57# Official speakers included in the Qwen3-TTS CustomVoice checkpoints.
 58_AVAILABLE_VOICES = [
 59    {
 60        "name": "Vivian",
 61        "native_language": "Chinese",
 62        "description": "Bright, slightly edgy young female voice",
 63    },
 64    {
 65        "name": "Serena",
 66        "native_language": "Chinese",
 67        "description": "Warm, gentle young female voice",
 68    },
 69    {
 70        "name": "Uncle_Fu",
 71        "native_language": "Chinese",
 72        "description": "Seasoned male voice with a low, mellow timbre",
 73    },
 74    {
 75        "name": "Dylan",
 76        "native_language": "Chinese",
 77        "description": "Youthful Beijing male voice with a clear timbre",
 78    },
 79    {
 80        "name": "Eric",
 81        "native_language": "Chinese",
 82        "description": "Lively Chengdu male voice with husky brightness",
 83    },
 84    {
 85        "name": "Ryan",
 86        "native_language": "English",
 87        "description": "Dynamic male voice with strong rhythmic drive",
 88    },
 89    {
 90        "name": "Aiden",
 91        "native_language": "English",
 92        "description": "Sunny American male voice with a clear midrange",
 93    },
 94    {
 95        "name": "Ono_Anna",
 96        "native_language": "Japanese",
 97        "description": "Playful Japanese female voice with a light timbre",
 98    },
 99    {
100        "name": "Sohee",
101        "native_language": "Korean",
102        "description": "Warm Korean female voice with rich emotion",
103    },
104]
105
106
107_MODEL_CACHE: dict[tuple[str, str, str], Qwen3TTSModel] = {}
108
109
110def _resolve_device(device: str | None = DEFAULT_DEVICE) -> str:
111    """デバイス設定を解決し、利用可能なCUDAデバイスまたはCPUを決定します。
112
113    ``auto`` が指定された場合、CUDAが利用可能であれば ``cuda:0`` を、
114    そうでなければ ``cpu`` を返します。
115
116    :param device: 使用するデバイス名(例: "cuda:0", "cpu", "auto")。
117                   Noneまたは"auto"の場合、自動的に決定されます。
118    :type device: str | None
119    :returns: 解決されたデバイス名。
120    :rtype: str
121    """
122    if device is None or str(device).strip().lower() == "auto":
123        return "cuda:0" if torch.cuda.is_available() else "cpu"
124    return str(device)
125
126
127def _resolve_dtype(device: str, dtype: str | torch.dtype | None) -> torch.dtype:
128    """データ型名をTorchのデータ型に解決します。
129
130    指定されたデバイスに適したデータ型を決定します。
131    ``auto`` が指定された場合、CUDAデバイスでは ``bfloat16`` または ``float16`` を、
132    CPUでは ``float32`` を使用します。
133
134    :param device: ターゲットデバイス名。
135    :type device: str
136    :param dtype: 使用するデータ型名(例: "auto", "bfloat16", "float16", "float32")
137                  または ``torch.dtype`` オブジェクト。
138    :type dtype: str | torch.dtype | None
139    :returns: 解決された ``torch.dtype`` オブジェクト。
140    :rtype: torch.dtype
141    :raises ValueError: サポートされていないデータ型が指定された場合。
142    """
143    if isinstance(dtype, torch.dtype):
144        return dtype
145
146    name = "auto" if dtype is None else str(dtype).strip().lower()
147    if name == "auto":
148        if device.startswith("cuda"):
149            return torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
150        return torch.float32
151
152    dtype_map = {
153        "bfloat16": torch.bfloat16,
154        "bf16": torch.bfloat16,
155        "float16": torch.float16,
156        "fp16": torch.float16,
157        "float32": torch.float32,
158        "fp32": torch.float32,
159    }
160    if name not in dtype_map:
161        raise ValueError(f"Unsupported dtype: {dtype}")
162    return dtype_map[name]
163
164
165def load_model(
166    model_id: str = DEFAULT_QWEN3_MODEL,
167    device: str = DEFAULT_DEVICE,
168    dtype: str | torch.dtype = DEFAULT_DTYPE,
169    *,
170    force_reload: bool = False,
171) -> Qwen3TTSModel:
172    """Qwen3-TTSモデルをロードし、キャッシュします。
173
174    FlashAttentionは意図的に要求されていません。これにより、個別にCUDA SDKをインストールせずに
175    Windows上でも通常のPyTorchアテンションパスが動作します。
176    モデルは、``model_id``, ``device``, ``dtype`` の組み合わせをキーとしてキャッシュされます。
177    既にキャッシュされているモデルは再ロードされません。
178
179    :param model_id: ロードするQwen3-TTSモデルのID(例: "Qwen/Qwen3-TTS-12Hz-0.6B-CustomVoice")。
180    :type model_id: str
181    :param device: モデルをロードするデバイス(例: "cuda:0", "cpu", "auto")。
182    :type device: str
183    :param dtype: モデルのデータ型(例: "auto", "bfloat16", "float16", "float32")。
184    :type dtype: str | torch.dtype
185    :param force_reload: Trueの場合、キャッシュを無視してモデルを強制的に再ロードします。
186    :type force_reload: bool
187    :returns: ロードされたQwen3TTSModelインスタンス。
188    :rtype: Qwen3TTSModel
189    """
190    resolved_device = _resolve_device(device)
191    resolved_dtype = _resolve_dtype(resolved_device, dtype)
192    cache_key = (model_id, resolved_device, str(resolved_dtype))
193
194    if force_reload:
195        _MODEL_CACHE.pop(cache_key, None)
196
197    if cache_key not in _MODEL_CACHE:
198        print(
199            "tktts_qwen3.load_model(): "
200            f"model={model_id}, device={resolved_device}, dtype={resolved_dtype}"
201        )
202        _MODEL_CACHE[cache_key] = Qwen3TTSModel.from_pretrained(
203            model_id,
204            device_map=resolved_device,
205            dtype=resolved_dtype,
206        )
207
208    return _MODEL_CACHE[cache_key]
209
210
211def unload_models() -> None:
212    """このモジュールによってキャッシュされたすべてのモデルを解放します。
213
214    GPUメモリを使用している場合、TorchのCUDAキャッシュもクリアします。
215
216    :returns: なし。
217    :rtype: None
218    """
219    _MODEL_CACHE.clear()
220    if torch.cuda.is_available():
221        torch.cuda.empty_cache()
222
223
224def get_available_voices_info(model: Qwen3TTSModel | None = None):
225    """利用可能な話者名とその説明を返します。
226
227    モデルインスタンスが提供された場合、そのモデルが報告する話者リストが使用されます。
228    これにより、将来のCustomVoiceチェックポイントとの互換性が保たれます。
229    モデルが提供されない場合、または ``get_supported_speakers`` メソッドがない場合は、
230    このモジュールにハードコードされたデフォルトのリストが使用されます。
231
232    :param model: 問い合わせるQwen3TTSModelインスタンス。Noneの場合はデフォルトリストを使用。
233    :type model: Qwen3TTSModel | None
234    :returns: 各話者の名前、母国語、説明を含む辞書のリスト。
235    :rtype: list[dict]
236    """
237    if model is None:
238        return [voice.copy() for voice in _AVAILABLE_VOICES]
239
240    get_speakers = getattr(model, "get_supported_speakers", None)
241    if not callable(get_speakers):
242        return [voice.copy() for voice in _AVAILABLE_VOICES]
243
244    known = {voice["name"].lower(): voice for voice in _AVAILABLE_VOICES}
245    voices = []
246    for speaker in get_speakers():
247        info = known.get(str(speaker).lower(), {})
248        voices.append(
249            {
250                "name": str(speaker),
251                "native_language": info.get("native_language", ""),
252                "description": info.get("description", ""),
253            }
254        )
255    return voices
256
257
258def get_available_voices(model: Qwen3TTSModel | None = None):
259    """利用可能な話者名をリストで返します。
260
261    これは ``get_available_voices_info()`` から話者名のみを抽出するヘルパー関数です。
262
263    :param model: 問い合わせるQwen3TTSModelインスタンス。Noneの場合はデフォルトリストを使用。
264    :type model: Qwen3TTSModel | None
265    :returns: 利用可能な話者名のリスト。
266    :rtype: list[str]
267    """
268    return [voice["name"] for voice in get_available_voices_info(model)]
269
270
271def list_available_voices(model: Qwen3TTSModel | None = None) -> bool:
272    """利用可能な話者をコンソールに出力します。
273
274    :param model: 問い合わせるQwen3TTSModelインスタンス。Noneの場合はデフォルトリストを使用。
275    :type model: Qwen3TTSModel | None
276    :returns: 利用可能な話者が出力された場合はTrue、そうでなければFalse。
277    :rtype: bool
278    """
279    print(f"=== 利用可能な {TTS_ENGINE_NAME} voices ===")
280    voices = get_available_voices_info(model)
281    if not voices:
282        return False
283
284    for voice in voices:
285        language = voice.get("native_language") or "-"
286        description = voice.get("description") or "-"
287        print(
288            f"  Name: {voice['name']}, "
289            f"Native language: {language}, Description: {description}"
290        )
291    return True
292
293
294def _speaker_key(name: str) -> str:
295    """話者名を正規化し、検索用のキーを生成します。
296
297    スペース、アンダースコア、ハイフン、括弧を削除し、すべて小文字に変換します。
298
299    :param name: 処理する話者名。
300    :type name: str
301    :returns: 正規化された話者キー。
302    :rtype: str
303    """
304    normalized = normalize_speaker(str(name))
305    return re.sub(r"[\s_\-()()]+", "", normalized).lower()
306
307
308def resolve_speaker_id(
309    speaker_name: str,
310    voices_info=None,
311    model: Qwen3TTSModel | None = None,
312):
313    """フルまたは部分的な話者名をQwenの話者名に解決します。
314
315    指定された話者名に基づいて、最も一致する話者名を検索します。
316    まず正規化された完全一致を優先し、次に部分一致を探します。
317
318    :param speaker_name: 解決したい話者名(フル名または部分名)。
319    :type speaker_name: str
320    :param voices_info: 利用可能な話者情報のリスト。Noneの場合、 ``get_available_voices_info`` を呼び出します。
321    :type voices_info: list[dict] | None
322    :param model: 話者情報を取得するために使用されるQwen3TTSModelインスタンス。
323                  ``voices_info`` が提供されている場合は使用されません。
324    :type model: Qwen3TTSModel | None
325    :returns: (更新されたvoices_info, 解決された話者名) のタプル。
326    :rtype: tuple[list[dict], str]
327    :raises ValueError: 話者名が空であるか、見つからない場合。
328    """
329    if voices_info is None:
330        voices_info = get_available_voices_info(model)
331
332    query = _speaker_key(speaker_name)
333    if not query:
334        raise ValueError("話者名が空です")
335
336    # Prefer an exact normalized match before accepting a partial match.
337    for voice in voices_info:
338        if query == _speaker_key(voice["name"]):
339            return voices_info, voice["name"]
340    for voice in voices_info:
341        if query in _speaker_key(voice["name"]):
342            return voices_info, voice["name"]
343
344    raise ValueError(
345        "❌ Error in tktts_qwen3.resolve_speaker_id(): "
346        f"話者 [{speaker_name}] が見つかりませんでした"
347    )
348
349
350def speak(
351    outfile,
352    text,
353    voice=DEFAULT_QWEN3_VOICE,
354    speak_rate=None,
355    speak_pitch=None,
356    *,
357    language: str = DEFAULT_LANGUAGE,
358    model: Qwen3TTSModel | None = None,
359    model_id: str = DEFAULT_QWEN3_MODEL,
360    device: str = DEFAULT_DEVICE,
361    dtype: str | torch.dtype = DEFAULT_DTYPE,
362    instruct: str | None = None,
363):
364    """Qwen3-TTS CustomVoiceモデルを使用して単一の音声ファイルを生成します。
365
366    ``speak_rate`` と ``speak_pitch`` は ``tktts_voicevox`` とのインターフェース互換性のために
367    残されていますが、CustomVoiceモデルはこれらのパラメータを公開していないため、
368    非デフォルト値は現在無視されます。
369
370    :param outfile: 生成された音声ファイルを保存するパス。
371    :type outfile: str | PathLike
372    :param text: 読み上げるテキスト。
373    :type text: str
374    :param voice: 使用する話者名。
375    :type voice: str
376    :param speak_rate: 音声の速さ(現在は無視されます)。
377    :type speak_rate: float | None
378    :param speak_pitch: 音声のピッチ(現在は無視されます)。
379    :type speak_pitch: float | None
380    :param language: 読み上げに使用する言語(例: "Japanese")。
381    :type language: str
382    :param model: 既存のQwen3TTSModelインスタンス。Noneの場合、 ``load_model`` を呼び出します。
383    :type model: Qwen3TTSModel | None
384    :param model_id: ``model`` がNoneの場合にロードするモデルID。
385    :type model_id: str
386    :param device: ``model`` がNoneの場合にモデルをロードするデバイス。
387    :type device: str
388    :param dtype: ``model`` がNoneの場合にモデルをロードするデータ型。
389    :type dtype: str | torch.dtype
390    :param instruct: 音声生成のための追加の指示テキスト。
391    :type instruct: str | None
392    :returns: 生成されたファイルのパス、または失敗した場合はNone。
393    :rtype: str | None
394    """
395    text = str(text).strip()
396    if not text:
397        print("❌ tktts_qwen3.speak(): 読み上げテキストが空です")
398        return None
399
400    if speak_rate is not None or speak_pitch is not None:
401        print(
402            "    ** Warning: Qwen3-TTS CustomVoiceでは "
403            "speak_rate/speak_pitchを直接指定できないため無視します"
404        )
405
406    if model is None:
407        model = load_model(model_id=model_id, device=device, dtype=dtype)
408
409    kwargs: dict[str, Any] = {
410        "text": text,
411        "language": language,
412        "speaker": voice,
413    }
414    if instruct:
415        kwargs["instruct"] = instruct
416
417    try:
418        with torch.inference_mode():
419            wavs, sample_rate = model.generate_custom_voice(**kwargs)
420    except Exception as exc:
421        print(f"❌ Qwen3-TTS生成エラー: {exc}")
422        return None
423
424    output_path = Path(outfile)
425    output_path.parent.mkdir(parents=True, exist_ok=True)
426    sf.write(str(output_path), wavs[0], sample_rate)
427
428    if output_path.exists():
429        print(f"    ** 一時ファイル [{output_path}] を保存しました")
430        return str(output_path)
431
432    print(f"    ** Error: ファイル [{output_path}] の出力に失敗しました")
433    return None
434
435
436def speak_dialogue(
437    dialogue,
438    replacements,
439    target_voices,
440    speakers=None,
441    temp_dir=".",
442    outfile=None,
443    ext="wav",
444    cfg=None,
445    *,
446    language: str = DEFAULT_LANGUAGE,
447    model: Qwen3TTSModel | None = None,
448    model_id: str = DEFAULT_QWEN3_MODEL,
449    device: str = DEFAULT_DEVICE,
450    dtype: str | torch.dtype = DEFAULT_DTYPE,
451):
452    """tktts対話シーケンス用の一時オーディオファイルを生成します。
453
454    この関数は、与えられた対話リストを処理し、各セグメントに対して個別の音声ファイルを生成します。
455    ``outfile`` は ``tktts_voicevox`` との互換性のために署名に残されていますが、この関数では使用されません。
456
457    :param dialogue: 対話アイテムのリスト。各アイテムはテキストまたは話者とテキストの組み合わせ。
458    :type dialogue: list
459    :param replacements: テキストに適用される置換規則の辞書。
460    :type replacements: dict
461    :param target_voices: 対話で話す話者の名前、または話者のリスト。
462    :type target_voices: str | list
463    :param speakers: 話者名と話者IDのマッピング辞書(オプション)。
464    :type speakers: dict | None
465    :param temp_dir: 一時音声ファイルを保存するディレクトリ。
466    :type temp_dir: str
467    :param outfile: この関数では使用されません(互換性のため残されています)。
468    :type outfile: Any
469    :param ext: 生成される音声ファイルの拡張子(例: "wav")。
470    :type ext: str
471    :param cfg: 設定オブジェクト( ``fspeak_rate``, ``fspeak_pitch``, ``qwen3_instruct``, ``monologue`` などの属性を含む場合があります)。
472    :type cfg: Any | None
473    :param language: 音声生成に使用する言語(例: "Japanese")。
474    :type language: str
475    :param model: 既存のQwen3TTSModelインスタンス。Noneの場合、 ``load_model`` を呼び出します。
476    :type model: Qwen3TTSModel | None
477    :param model_id: ``model`` がNoneの場合にロードするモデルID。
478    :type model_id: str
479    :param device: ``model`` がNoneの場合にモデルをロードするデバイス。
480    :type device: str
481    :param dtype: ``model`` がNoneの場合にモデルをロードするデータ型。
482    :type dtype: str | torch.dtype
483    :returns: (成功したかどうかのブール値, 生成された一時ファイルのパスのリスト) のタプル。
484    :rtype: tuple[bool, list[str]]
485    """
486    del outfile  # Kept in the signature for compatibility with tktts_voicevox.
487    speakers = {} if speakers is None else speakers
488
489    print("tktts_qwen3.speak_dialogue(): target_voices:", target_voices)
490
491    if model is None:
492        try:
493            model = load_model(model_id=model_id, device=device, dtype=dtype)
494        except Exception as exc:
495            print(f"❌ Qwen3-TTSモデル読み込みエラー: {exc}")
496            return False, []
497
498    tmpfiles = []
499    voices_info = get_available_voices_info(model)
500    idx = 1
501    is_monologue = bool(getattr(cfg, "monologue", False))
502
503    for i, dialogue_item in enumerate(dialogue):
504        print()
505        print(f"Dialogue {i:04d}:")
506        dialogue_list = split_dialogue(
507            dialogue_item,
508            target_voices,
509            speakers=speakers,
510            default_voice=DEFAULT_QWEN3_VOICE,
511            is_monologue=is_monologue,
512        )
513
514        for speaker, text in dialogue_list:
515            tmpfile = os.path.join(temp_dir, f"tmp_{idx:03d}.{ext}")
516            text = apply_replacements(text, replacements)
517            if isinstance(target_voices, str):
518                speaker = target_voices
519
520            try:
521                voices_info, target_voice = resolve_speaker_id(
522                    speaker,
523                    voices_info=voices_info,
524                    model=model,
525                )
526            except ValueError as exc:
527                print(exc)
528                return False, tmpfiles
529
530            print(f"  {idx:04d}: voice={speaker} ({target_voice}): {text}")
531
532            speak_rate = getattr(cfg, "fspeak_rate", None)
533            speak_pitch = getattr(cfg, "fspeak_pitch", None)
534            instruct = getattr(cfg, "qwen3_instruct", None)
535
536            generated = speak(
537                outfile=tmpfile,
538                text=text,
539                voice=target_voice,
540                speak_rate=speak_rate,
541                speak_pitch=speak_pitch,
542                language=language,
543                model=model,
544                instruct=instruct,
545            )
546            if generated is None:
547                return False, tmpfiles
548
549            tmpfiles.append(tmpfile)
550            idx += 1
551
552    return True, tmpfiles
553
554
555if __name__ == "__main__":
556    list_available_voices()