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

1""" 

2Wrapper for Kokoro text-to-speech model. 

3""" 

4import logging 

5import tempfile 

6import asyncio 

7 

8import torch 

9 

10from torch import inference_mode 

11from torch import Tensor 

12 

13from typing import override 

14from typing import List 

15from typing import Dict 

16from typing import Any 

17from typing import Union 

18from typing import Optional 

19 

20from enum import Enum 

21 

22from kokoro import KPipeline 

23 

24from model_timing import GenTimer 

25from wrapper_model import ModelGeneration 

26from media_utils import save_audio 

27 

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} 

42 

43 

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" 

58 

59 

60class KokoroGeneration(ModelGeneration): 

61 """Handle audio generation using the Kokoro model.""" 

62 

63 def __init__(self) -> None: 

64 super().__init__("kokoro") 

65 

66 # Model components: language -> model pipeline 

67 self.kokoro: Dict[str, KPipeline] = {} 

68 

69 def __del__(self) -> None: 

70 if self.kokoro: 

71 self.kokoro = {} # Release model pipelines 

72 

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") 

87 

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") 

105 

106 def init_model_parallelism(self) -> None: 

107 if self.world_size > 1: 

108 logging.warning("Kokoro does not support distributed parallelism.") 

109 

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") 

121 

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") 

125 

126 job_id = data_json.get("job_id", None) 

127 

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 } 

145 

146 @torch.inference_mode() 

147 async def warmup(self) -> None: 

148 logging.info("Warmup for Kokoro generation.") 

149 await self.generate(text="Warmup") 

150 

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) 

163 

164 self.running = True # We can run in parallel but good to know if we are running 

165 

166 audios = [] 

167 try: 

168 # Clean up text 

169 text = text.replace("*", "") 

170 

171 if lang_code not in self.kokoro: 

172 raise ValueError(f"Unsupported language code: {lang_code}") 

173 

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") 

182 

183 # text, phonemes, audio = gs, ps, audio 

184 for gs, ps, audio in audio_generator: 

185 audios.append(audio) 

186 

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") 

195 

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 

210 

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" 

215 

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") 

223 

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