codemo/fish-speech-1
0
1import base642import io3import json4import queue5import random6import sys7import traceback8import wave9from argparse import ArgumentParser10from http import HTTPStatus11from pathlib import Path12from typing import Annotated, Any, Literal, Optional13 14import numpy as np15import ormsgpack16import pyrootutils17import soundfile as sf18import torch19import torchaudio20from baize.datastructures import ContentType21from kui.asgi import (22 Body,23 FactoryClass,24 HTTPException,25 HttpRequest,26 HttpView,27 JSONResponse,28 Kui,29 OpenAPI,30 StreamResponse,31)32from kui.asgi.routing import MultimethodRoutes33from loguru import logger34from pydantic import BaseModel, Field, conint35 36pyrootutils.setup_root(__file__, indicator=".project-root", pythonpath=True)37 38# from fish_speech.models.vqgan.lit_module import VQGAN39from fish_speech.models.vqgan.modules.firefly import FireflyArchitecture40from fish_speech.text.chn_text_norm.text import Text as ChnNormedText41from fish_speech.utils import autocast_exclude_mps42from tools.commons import ServeReferenceAudio, ServeTTSRequest43from tools.file import AUDIO_EXTENSIONS, audio_to_bytes, list_files, read_ref_text44from tools.llama.generate import (45 GenerateRequest,46 GenerateResponse,47 WrappedGenerateResponse,48 launch_thread_safe_queue,49)50from tools.vqgan.inference import load_model as load_decoder_model51 52 53def wav_chunk_header(sample_rate=44100, bit_depth=16, channels=1):54 buffer = io.BytesIO()55 56 with wave.open(buffer, "wb") as wav_file:57 wav_file.setnchannels(channels)58 wav_file.setsampwidth(bit_depth // 8)59 wav_file.setframerate(sample_rate)60 61 wav_header_bytes = buffer.getvalue()62 buffer.close()63 return wav_header_bytes64 65 66# Define utils for web server67async def http_execption_handler(exc: HTTPException):68 return JSONResponse(69 dict(70 statusCode=exc.status_code,71 message=exc.content,72 error=HTTPStatus(exc.status_code).phrase,73 ),74 exc.status_code,75 exc.headers,76 )77 78 79async def other_exception_handler(exc: "Exception"):80 traceback.print_exc()81 82 status = HTTPStatus.INTERNAL_SERVER_ERROR83 return JSONResponse(84 dict(statusCode=status, message=str(exc), error=status.phrase),85 status,86 )87 88 89def load_audio(reference_audio, sr):90 if len(reference_audio) > 255 or not Path(reference_audio).exists():91 audio_data = reference_audio92 reference_audio = io.BytesIO(audio_data)93 94 waveform, original_sr = torchaudio.load(95 reference_audio, backend="soundfile" if sys.platform == "linux" else "soundfile"96 )97 98 if waveform.shape[0] > 1:99 waveform = torch.mean(waveform, dim=0, keepdim=True)100 101 if original_sr != sr:102 resampler = torchaudio.transforms.Resample(orig_freq=original_sr, new_freq=sr)103 waveform = resampler(waveform)104 105 audio = waveform.squeeze().numpy()106 return audio107 108 109def encode_reference(*, decoder_model, reference_audio, enable_reference_audio):110 if enable_reference_audio and reference_audio is not None:111 # Load audios, and prepare basic info here112 reference_audio_content = load_audio(113 reference_audio, decoder_model.spec_transform.sample_rate114 )115 116 audios = torch.from_numpy(reference_audio_content).to(decoder_model.device)[117 None, None, :118 ]119 audio_lengths = torch.tensor(120 [audios.shape[2]], device=decoder_model.device, dtype=torch.long121 )122 logger.info(123 f"Loaded audio with {audios.shape[2] / decoder_model.spec_transform.sample_rate:.2f} seconds"124 )125 126 # VQ Encoder127 if isinstance(decoder_model, FireflyArchitecture):128 prompt_tokens = decoder_model.encode(audios, audio_lengths)[0][0]129 130 logger.info(f"Encoded prompt: {prompt_tokens.shape}")131 else:132 prompt_tokens = None133 logger.info("No reference audio provided")134 135 return prompt_tokens136 137 138def decode_vq_tokens(139 *,140 decoder_model,141 codes,142):143 feature_lengths = torch.tensor([codes.shape[1]], device=decoder_model.device)144 logger.info(f"VQ features: {codes.shape}")145 146 if isinstance(decoder_model, FireflyArchitecture):147 # VQGAN Inference148 return decoder_model.decode(149 indices=codes[None],150 feature_lengths=feature_lengths,151 )[0].squeeze()152 153 raise ValueError(f"Unknown model type: {type(decoder_model)}")154 155 156routes = MultimethodRoutes(base_class=HttpView)157 158 159def get_content_type(audio_format):160 if audio_format == "wav":161 return "audio/wav"162 elif audio_format == "flac":163 return "audio/flac"164 elif audio_format == "mp3":165 return "audio/mpeg"166 else:167 return "application/octet-stream"168 169 170@torch.inference_mode()171def inference(req: ServeTTSRequest):172 173 idstr: str | None = req.reference_id174 if idstr is not None:175 ref_folder = Path("references") / idstr176 ref_folder.mkdir(parents=True, exist_ok=True)177 ref_audios = list_files(178 ref_folder, AUDIO_EXTENSIONS, recursive=True, sort=False179 )180 prompt_tokens = [181 encode_reference(182 decoder_model=decoder_model,183 reference_audio=audio_to_bytes(str(ref_audio)),184 enable_reference_audio=True,185 )186 for ref_audio in ref_audios187 ]188 prompt_texts = [189 read_ref_text(str(ref_audio.with_suffix(".lab")))190 for ref_audio in ref_audios191 ]192 193 else:194 # Parse reference audio aka prompt195 refs = req.references196 if refs is None:197 refs = []198 prompt_tokens = [199 encode_reference(200 decoder_model=decoder_model,201 reference_audio=ref.audio,202 enable_reference_audio=True,203 )204 for ref in refs205 ]206 prompt_texts = [ref.text for ref in refs]207 208 # LLAMA Inference209 request = dict(210 device=decoder_model.device,211 max_new_tokens=req.max_new_tokens,212 text=(213 req.text214 if not req.normalize215 else ChnNormedText(raw_text=req.text).normalize()216 ),217 top_p=req.top_p,218 repetition_penalty=req.repetition_penalty,219 temperature=req.temperature,220 compile=args.compile,221 iterative_prompt=req.chunk_length > 0,222 chunk_length=req.chunk_length,223 max_length=2048,224 prompt_tokens=prompt_tokens,225 prompt_text=prompt_texts,226 )227 228 response_queue = queue.Queue()229 llama_queue.put(230 GenerateRequest(231 request=request,232 response_queue=response_queue,233 )234 )235 236 if req.streaming:237 yield wav_chunk_header()238 239 segments = []240 while True:241 result: WrappedGenerateResponse = response_queue.get()242 if result.status == "error":243 raise result.response244 break245 246 result: GenerateResponse = result.response247 if result.action == "next":248 break249 250 with autocast_exclude_mps(251 device_type=decoder_model.device.type, dtype=args.precision252 ):253 fake_audios = decode_vq_tokens(254 decoder_model=decoder_model,255 codes=result.codes,256 )257 258 fake_audios = fake_audios.float().cpu().numpy()259 260 if req.streaming:261 yield (fake_audios * 32768).astype(np.int16).tobytes()262 else:263 segments.append(fake_audios)264 265 if req.streaming:266 return267 268 if len(segments) == 0:269 raise HTTPException(270 HTTPStatus.INTERNAL_SERVER_ERROR,271 content="No audio generated, please check the input text.",272 )273 274 fake_audios = np.concatenate(segments, axis=0)275 yield fake_audios276 277 278async def inference_async(req: ServeTTSRequest):279 for chunk in inference(req):280 yield chunk281 282 283async def buffer_to_async_generator(buffer):284 yield buffer285 286 287@routes.http.post("/v1/tts")288async def api_invoke_model(289 req: Annotated[ServeTTSRequest, Body(exclusive=True)],290):291 """292 Invoke model and generate audio293 """294 295 if args.max_text_length > 0 and len(req.text) > args.max_text_length:296 raise HTTPException(297 HTTPStatus.BAD_REQUEST,298 content=f"Text is too long, max length is {args.max_text_length}",299 )300 301 if req.streaming and req.format != "wav":302 raise HTTPException(303 HTTPStatus.BAD_REQUEST,304 content="Streaming only supports WAV format",305 )306 307 if req.streaming:308 return StreamResponse(309 iterable=inference_async(req),310 headers={311 "Content-Disposition": f"attachment; filename=audio.{req.format}",312 },313 content_type=get_content_type(req.format),314 )315 else:316 fake_audios = next(inference(req))317 buffer = io.BytesIO()318 sf.write(319 buffer,320 fake_audios,321 decoder_model.spec_transform.sample_rate,322 format=req.format,323 )324 325 return StreamResponse(326 iterable=buffer_to_async_generator(buffer.getvalue()),327 headers={328 "Content-Disposition": f"attachment; filename=audio.{req.format}",329 },330 content_type=get_content_type(req.format),331 )332 333 334@routes.http.post("/v1/health")335async def api_health():336 """337 Health check338 """339 340 return JSONResponse({"status": "ok"})341 342 343def parse_args():344 parser = ArgumentParser()345 parser.add_argument(346 "--llama-checkpoint-path",347 type=str,348 default="checkpoints/fish-speech-1.4",349 )350 parser.add_argument(351 "--decoder-checkpoint-path",352 type=str,353 default="checkpoints/fish-speech-1.4/firefly-gan-vq-fsq-8x1024-21hz-generator.pth",354 )355 parser.add_argument("--decoder-config-name", type=str, default="firefly_gan_vq")356 parser.add_argument("--device", type=str, default="cuda")357 parser.add_argument("--half", action="store_true")358 parser.add_argument("--compile", action="store_true")359 parser.add_argument("--max-text-length", type=int, default=0)360 parser.add_argument("--listen", type=str, default="127.0.0.1:8080")361 parser.add_argument("--workers", type=int, default=1)362 parser.add_argument("--use-auto-rerank", type=bool, default=True)363 364 return parser.parse_args()365 366 367# Define Kui app368openapi = OpenAPI(369 {370 "title": "Fish Speech API",371 },372).routes373 374 375class MsgPackRequest(HttpRequest):376 async def data(self) -> Annotated[Any, ContentType("application/msgpack")]:377 if self.content_type == "application/msgpack":378 return ormsgpack.unpackb(await self.body)379 380 raise HTTPException(381 HTTPStatus.UNSUPPORTED_MEDIA_TYPE,382 headers={"Accept": "application/msgpack"},383 )384 385 386app = Kui(387 routes=routes + openapi[1:], # Remove the default route388 exception_handlers={389 HTTPException: http_execption_handler,390 Exception: other_exception_handler,391 },392 factory_class=FactoryClass(http=MsgPackRequest),393 cors_config={},394)395 396 397if __name__ == "__main__":398 399 import uvicorn400 401 args = parse_args()402 args.precision = torch.half if args.half else torch.bfloat16403 404 logger.info("Loading Llama model...")405 llama_queue = launch_thread_safe_queue(406 checkpoint_path=args.llama_checkpoint_path,407 device=args.device,408 precision=args.precision,409 compile=args.compile,410 )411 logger.info("Llama model loaded, loading VQ-GAN model...")412 413 decoder_model = load_decoder_model(414 config_name=args.decoder_config_name,415 checkpoint_path=args.decoder_checkpoint_path,416 device=args.device,417 )418 419 logger.info("VQ-GAN model loaded, warming up...")420 421 # Dry run to check if the model is loaded correctly and avoid the first-time latency422 list(423 inference(424 ServeTTSRequest(425 text="Hello world.",426 references=[],427 reference_id=None,428 max_new_tokens=0,429 top_p=0.7,430 repetition_penalty=1.2,431 temperature=0.7,432 emotion=None,433 format="wav",434 )435 )436 )437 438 logger.info(f"Warming up done, starting server at http://{args.listen}")439 host, port = args.listen.split(":")440 uvicorn.run(app, host=host, port=int(port), workers=args.workers, log_level="info")441 