CoolFace
Apppublic

codemo/fish-speech-1

sourceHugging Facecc-by-nc-sa-4.0updated 2y agoView on Hugging Face
0likes
api.py441 linesDownload Raw Back to tools
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