Coverage for wrapper/kokoro/wrapper_kokoro.py: 87%
126 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1"""
2Wrapper for Kokoro text-to-speech model.
3"""
4import logging
5import tempfile
6import asyncio
8import torch
10from torch import inference_mode
11from torch import Tensor
13from typing import override
14from typing import List
15from typing import Dict
16from typing import Any
17from typing import Union
18from typing import Optional
20from enum import Enum
22from kokoro import KPipeline
24from model_timing import GenTimer
25from wrapper_model import ModelGeneration
26from media_utils import save_audio
28VOICES = {
29 "female": [
30 "af_heart",
31 "af_bella",
32 "af_kore",
33 "af_nicole",
34 ],
35 "male": [
36 "am_adam",
37 "am_puck",
38 "am_michael",
39 "am_fenrir",
40 ]
41}
44class Language(str, Enum):
45 """
46 Supported languages.
47 https://github.com/hexgrad/kokoro/blob/main/kokoro/pipeline.py
48 """
49 AMERICAN_ENGLISH = "a"
50 BRITISH_ENGLISH = "b"
51 SPANISH = "e"
52 FRENCH = "f"
53 HINDI = "h"
54 ITALIAN = "i"
55 BRAZILIAN_PORTUGUESE = "p"
56 JAPANESE = "j"
57 MANDARIN_CHINESE = "z"
60class KokoroGeneration(ModelGeneration):
61 """Handle audio generation using the Kokoro model."""
63 def __init__(self) -> None:
64 super().__init__("kokoro")
66 # Model components: language -> model pipeline
67 self.kokoro: Dict[str, KPipeline] = {}
69 def __del__(self) -> None:
70 if self.kokoro:
71 self.kokoro = {} # Release model pipelines
73 def init_parallelism(self) -> None:
74 self.load_timer.start("torch_dist")
75 # No real parallelism as it runs with a single GPU or CPU
76 if torch.cuda.is_available():
77 self.rank = 0
78 self.local_rank = 0
79 self.world_size = 1
80 self.device_id: Union[int, str] = self.local_rank
81 self.device = torch.device(f"cuda:{self.device_id}")
82 torch.cuda.set_device(self.local_rank)
83 else:
84 self.device_id = "cpu"
85 self.device = torch.device(self.device_id)
86 self.load_timer.end("torch_dist")
88 def load_model(self) -> None:
89 self.load_timer.start("kokoro")
90 model_name = "hexgrad/Kokoro-82M"
91 for lang in Language:
92 lang_code = lang.value
93 try:
94 logging.info(f"Loading Kokoro model for language {lang} ({lang_code})")
95 self.load_timer.start(f"kokoro_{lang_code}")
96 kokoro = KPipeline(
97 repo_id=model_name,
98 device=self.device,
99 lang_code=lang_code)
100 self.kokoro[lang_code] = kokoro
101 self.load_timer.end(f"kokoro_{lang_code}")
102 except Exception as ex:
103 logging.error(f"Cannot load language {lang} ({lang_code}): {str(ex)}")
104 self.load_timer.end("kokoro")
106 def init_model_parallelism(self) -> None:
107 if self.world_size > 1:
108 logging.warning("Kokoro does not support distributed parallelism.")
110 def model_compile(self) -> None:
111 if not self.torch_compile:
112 return
113 self.load_timer.start("compile")
114 for lang in Language:
115 lang_code = lang.value
116 if lang_code in self.kokoro:
117 self.kokoro[lang_code] = torch.compile(
118 self.kokoro[lang_code],
119 mode="reduce-overhead")
120 self.load_timer.end("compile")
122 async def get_rest_args(self, data_json: dict) -> Dict[str, Any]:
123 if data_json is None:
124 raise ValueError("Missing JSON body")
126 job_id = data_json.get("job_id", None)
128 text = data_json.get("text", None)
129 if text is None:
130 raise ValueError("Missing 'text' parameter")
131 # https://github.com/hexgrad/kokoro/blob/main/kokoro.js/src/voices.js
132 voice = data_json.get("voice", "af_heart")
133 speed = float(data_json.get("speed", 1.0))
134 lang_code = data_json.get("lang_code", Language.AMERICAN_ENGLISH.value)
135 return {
136 "task": self.model_name,
137 "args": {
138 "job_id": job_id,
139 "text": text,
140 "voice": voice,
141 "speed": speed,
142 "lang_code": lang_code,
143 }
144 }
146 @torch.inference_mode()
147 async def warmup(self) -> None:
148 logging.info("Warmup for Kokoro generation.")
149 await self.generate(text="Warmup")
151 @override
152 @inference_mode()
153 async def generate(
154 self,
155 text: str,
156 voice: str = "af_heart",
157 speed: float = 1.0,
158 lang_code: str = Language.AMERICAN_ENGLISH.value,
159 job_id: Optional[str] = None,
160 output_type: str = "audio_path",
161 ) -> Optional[Union[str, Tensor]]:
162 gen_timer = self._new_gen_timer(job_id)
164 self.running = True # We can run in parallel but good to know if we are running
166 audios = []
167 try:
168 # Clean up text
169 text = text.replace("*", "")
171 if lang_code not in self.kokoro:
172 raise ValueError(f"Unsupported language code: {lang_code}")
174 gen_timer.start("kokoro")
175 audio_generator = await asyncio.to_thread(
176 self.kokoro[lang_code],
177 text=text,
178 voice=voice,
179 speed=speed
180 )
181 gen_timer.end("kokoro")
183 # text, phonemes, audio = gs, ps, audio
184 for gs, ps, audio in audio_generator:
185 audios.append(audio)
187 return await self._output_audio(
188 job_id=job_id,
189 gen_timer=gen_timer,
190 audios=audios,
191 output_type=output_type)
192 finally:
193 self.running = False
194 gen_timer.end("total")
196 async def _output_audio(
197 self,
198 job_id: Optional[str],
199 gen_timer: GenTimer,
200 audios: List[Tensor],
201 output_type: str = "audio_path", # "audio_path"
202 ) -> Optional[Union[str, Tensor]]:
203 gen_timer.start("output_audio")
204 if len(audios) > 1:
205 logging.warning("Multiple audio chunks generated, returning the first one.")
206 try:
207 for audio in audios:
208 if output_type == "tensor":
209 return audio
211 if not job_id:
212 audio_path = tempfile.NamedTemporaryFile(suffix=".wav", delete=False).name
213 else:
214 audio_path = f"/tmp/{job_id}.wav"
216 audio_path = save_audio(
217 audio=audio,
218 audio_path=audio_path)
219 return audio_path
220 return None
221 finally:
222 gen_timer.end("output_audio")
224 def get_health(self) -> Dict[str, Any]:
225 ret = super().get_health()
226 if torch.cuda.is_available():
227 ret["gpu"] = torch.cuda.get_device_name(0)
228 return ret