CoolFace
Apppublic

lenML/ChatTTS-Forge

sourceHugging Faceagpl-3.0updated 2y agoView on Hugging Face
301likes
speaker_api.py125 linesDownload Raw Back to impl
1import torch2from fastapi import HTTPException3from pydantic import BaseModel4 5from modules.api import utils as api_utils6from modules.api.Api import APIManager7from modules.speaker import speaker_mgr8 9 10class CreateSpeaker(BaseModel):11    name: str12    gender: str13    describe: str14    tensor: list = None15    seed: int = None16 17 18class UpdateSpeaker(BaseModel):19    id: str20    name: str21    gender: str22    describe: str23    tensor: list24 25 26class SpeakerDetail(BaseModel):27    id: str28    with_emb: bool = False29 30 31class SpeakersUpdate(BaseModel):32    speakers: list33 34 35def setup(app: APIManager):36 37    @app.get("/v1/speakers/list", response_model=api_utils.BaseResponse)38    async def list_speakers():39        return api_utils.success_response(40            [spk.to_json() for spk in speaker_mgr.list_speakers()]41        )42 43    @app.post("/v1/speakers/refresh", response_model=api_utils.BaseResponse)44    async def refresh_speakers():45        speaker_mgr.refresh_speakers()46        return api_utils.success_response(None)47 48    @app.post("/v1/speakers/update", response_model=api_utils.BaseResponse)49    async def update_speakers(request: SpeakersUpdate):50        for spk in request.speakers:51            speaker = speaker_mgr.get_speaker_by_id(spk["id"])52            if speaker is None:53                raise HTTPException(54                    status_code=404, detail=f"Speaker not found: {spk['id']}"55                )56            speaker.name = spk.get("name", speaker.name)57            speaker.gender = spk.get("gender", speaker.gender)58            speaker.describe = spk.get("describe", speaker.describe)59            if (60                spk.get("tensor")61                and isinstance(spk["tensor"], list)62                and len(spk["tensor"]) > 063            ):64                # number array => Tensor65                speaker.emb = torch.tensor(spk["tensor"])66        speaker_mgr.save_all()67 68        return api_utils.success_response(None)69 70    @app.post("/v1/speaker/create", response_model=api_utils.BaseResponse)71    async def create_speaker(request: CreateSpeaker):72        if (73            request.tensor74            and isinstance(request.tensor, list)75            and len(request.tensor) > 076        ):77            # from tensor78            tensor = torch.tensor(request.tensor)79            speaker = speaker_mgr.create_speaker_from_tensor(80                tensor=tensor,81                name=request.name,82                gender=request.gender,83                describe=request.describe,84            )85        elif request.seed:86            # from seed87            speaker = speaker_mgr.create_speaker_from_seed(88                seed=request.seed,89                name=request.name,90                gender=request.gender,91                describe=request.describe,92            )93        else:94            raise HTTPException(95                status_code=400, detail="Missing tensor or seed in request"96            )97        return api_utils.success_response(speaker.to_json())98 99    @app.post("/v1/speaker/update", response_model=api_utils.BaseResponse)100    async def update_speaker(request: UpdateSpeaker):101        speaker = speaker_mgr.get_speaker_by_id(request.id)102        if speaker is None:103            raise HTTPException(104                status_code=404, detail=f"Speaker not found: {request.id}"105            )106        speaker.name = request.name107        speaker.gender = request.gender108        speaker.describe = request.describe109        if (110            request.tensor111            and isinstance(request.tensor, list)112            and len(request.tensor) > 0113        ):114            # number array => Tensor115            speaker.emb = torch.tensor(request.tensor)116        speaker_mgr.update_speaker(speaker)117        return api_utils.success_response(None)118 119    @app.post("/v1/speaker/detail", response_model=api_utils.BaseResponse)120    async def speaker_detail(request: SpeakerDetail):121        speaker = speaker_mgr.get_speaker_by_id(request.id)122        if speaker is None:123            raise HTTPException(status_code=404, detail="Speaker not found")124        return api_utils.success_response(speaker.to_json(with_emb=request.with_emb))125