CoolFace
Apppublic

Govendas/diarization

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
main.py344 linesDownload Raw Back to app
1import asyncio2import base643import io4import os5import subprocess6import time7from typing import List, Optional, Tuple, Dict8 9from fastapi import FastAPI, File, UploadFile, HTTPException10from fastapi.responses import StreamingResponse, HTMLResponse, Response11from fastapi.staticfiles import StaticFiles12from pydantic import BaseModel, Field13 14from app.fastapi_custom import custom_openapi15 16# deps de áudio/diarização17from pydub import AudioSegment18from pyannote.audio import Pipeline19import torch20import soundfile as sf21import warnings22 23import tempfile, textwrap, subprocess24from pathlib import Path25from fastapi import HTTPException26 27warnings.filterwarnings("ignore", message="torchaudio._backend")28warnings.filterwarnings("ignore", message="AudioMetaData has been deprecated")29 30# ---------------------------31# Config / globals32# ---------------------------33HF_TOKEN = os.getenv("HF_TOKEN")34 35MAX_CONCURRENCY = int(os.getenv("MAX_THREADS", "1"))   # igual ao TTS: controle por env36SEG_BS = int(os.getenv("PYANNOTE_SEG_BS", "256"))37EMB_BS = int(os.getenv("PYANNOTE_EMB_BS", "256"))38GAP_TOL = float(os.getenv("GAP_TOL", "0.50"))39MIN_DUR = float(os.getenv("MIN_DUR", "0.20"))40MAX_BLOCK = float(os.getenv("MAX_BLOCK_SEC", "12.0"))41 42# ---------------------------43# FastAPI app44# ---------------------------45app = FastAPI(title="API de Diarização")46app.mount("/static", StaticFiles(directory="static"), name="static")47app.openapi = lambda: custom_openapi(app)48 49# fila + semáforo (mesma estratégia do TTS)50queue = asyncio.Queue()51semaphore = asyncio.Semaphore(MAX_CONCURRENCY)52 53# pipeline global (carregado no startup)54PIPELINE: Pipeline | None = None55 56# ---------------------------57# Models (request/response)58# ---------------------------59class DiarizeRequest(BaseModel):60    # áudio em base64 (qualquer formato que o ffmpeg leia: mp3, m4a, mp4, wav…)61    audio_base64: str = Field(..., description="Áudio em base64 (arquivo binário completo)")62    # parâmetros opcionais para controlar fusões (defaults por env)63    gap_tol: Optional[float] = Field(None, description="Gap máximo para fundir segmentos consecutivos do mesmo falante (s)")64    min_dur: Optional[float] = Field(None, description="Duração mínima de um segmento após fusão (s)")65    max_block_sec: Optional[float] = Field(None, description="Limite de duração por bloco do mesmo falante (s)")66 67class SegmentItem(BaseModel):68    start: float69    end: float70    speaker: int71    audio_base64: str  # WAV base64 do trecho72 73class DiarizeResponse(BaseModel):74    segments: List[SegmentItem]75    num_speakers: int76    # opcional: métrica de tempo77    preprocess_sec: float78    diarization_sec: float79    slicing_sec: float80    total_sec: float81 82# ---------------------------83# Utilitários84# ---------------------------85def ffmpeg_bytes_to_wav16k(audio_bytes: bytes) -> bytes:86    if not audio_bytes:87        raise HTTPException(status_code=400, detail="Arquivo de áudio vazio.")88 89    # 1) grava entrada num arquivo temp com extensão “genérica”90    tmpdir = tempfile.mkdtemp(prefix="dia_in_")91    in_path = str(Path(tmpdir) / "input_media")92    out_path = str(Path(tmpdir) / "out.wav")93    with open(in_path, "wb") as f:94        f.write(audio_bytes)95 96    # 2) roda ffmpeg (arquivo → arquivo) e depois lê bytes97    cmd = [98        "ffmpeg", "-hide_banner", "-loglevel", "error", "-nostdin", "-y",99        "-threads", "0",100        "-i", in_path,101        "-ac", "1", "-ar", "16000",102        "-vn", "-sn", "-dn",103        "-c:a", "pcm_s16le",104        out_path105    ]106    try:107        res = subprocess.run(cmd, capture_output=True, check=False)108        if res.returncode != 0 or not Path(out_path).exists():109            err = (res.stderr or b"").decode("utf-8", errors="ignore")110            err_short = "\n".join(err.splitlines()[-6:])  # últimas linhas111            raise HTTPException(112                status_code=400,113                detail="Falha ao converter áudio com ffmpeg.\n" + textwrap.shorten(err_short, width=600)114            )115        return Path(out_path).read_bytes()116    except FileNotFoundError:117        # ffmpeg não instalado118        raise HTTPException(status_code=500, detail="ffmpeg não encontrado no PATH.")119 120def merge_adjacent_same_speaker(segs: List[Tuple[float,float,str]], gap_tol: float, min_dur: float):121    """Funde segmentos consecutivos do mesmo speaker se o gap <= gap_tol."""122    if not segs: return []123    merged = []124    cur_s, cur_e, cur_spk = segs[0]125    for s, e, spk in segs[1:]:126        if spk == cur_spk and s - cur_e <= gap_tol:127            cur_e = max(cur_e, e)128        else:129            if cur_e - cur_s >= min_dur:130                merged.append((cur_s, cur_e, cur_spk))131            cur_s, cur_e, cur_spk = s, e, spk132    if cur_e - cur_s >= min_dur:133        merged.append((cur_s, cur_e, cur_spk))134    return merged135 136def fuse_by_duration(segs: List[Tuple[float,float,str]], max_block_sec: float):137    """Funde trechos contíguos do MESMO speaker até um limite de duração."""138    if not segs: return []139    out = []140    cur_s, cur_e, cur_spk = segs[0]141    for s, e, spk in segs[1:]:142        if spk == cur_spk and (e - cur_s) <= max_block_sec:143            cur_e = e144        else:145            out.append((cur_s, cur_e, cur_spk))146            cur_s, cur_e, cur_spk = s, e, spk147    out.append((cur_s, cur_e, cur_spk))148    return out149 150def run_diarization_on_wav_bytes(wav_bytes: bytes) -> List[Tuple[float,float,str]]:151    """Roda o pipeline global (pyannote) e devolve (start, end, 'SPEAKER_xx')."""152    assert PIPELINE is not None, "Pipeline não inicializado."153    # pyannote aceita caminho ou dict; aqui vamos pelo caminho temp154    import tempfile155    from pathlib import Path156    tmpdir = tempfile.mkdtemp(prefix="dia_")157    wav_path = str(Path(tmpdir) / "in.wav")158    with open(wav_path, "wb") as f:159        f.write(wav_bytes)160 161    with torch.inference_mode():162        diar = PIPELINE(wav_path)163 164    segs: List[Tuple[float,float,str]] = []165    for turn, _, spk in diar.itertracks(yield_label=True):166        segs.append((float(turn.start), float(turn.end), spk))167    segs.sort(key=lambda x: x[0])168    return segs169 170def slice_segments_to_wav_base64(wav_bytes: bytes, segs: List[Tuple[float,float,str]]) -> Tuple[List[SegmentItem], int]:171    """Corta o WAV em chunks pelos segmentos e retorna lista de SegmentItem + #speakers."""172    audio = AudioSegment.from_file(io.BytesIO(wav_bytes), format="wav")173    # mapeia SPEAKER_00 -> 1, SPEAKER_01 -> 2…174    speaker_map: Dict[str,int] = {}175    next_id = 1176    items: List[SegmentItem] = []177    for s, e, spk in segs:178        if spk not in speaker_map:179            speaker_map[spk] = next_id; next_id += 1180        seg_audio = audio[int(s*1000):int(e*1000)]181        buf = io.BytesIO()182        # salva WAV (16k mono já está)183        seg_audio.export(buf, format="wav")184        buf.seek(0)185        b64 = base64.b64encode(buf.read()).decode("utf-8")186        items.append(SegmentItem(187            start=s, end=e, speaker=speaker_map[spk], audio_base64=b64188        ))189    return items, len(speaker_map)190 191# ---------------------------192# Worker assíncrono193# ---------------------------194async def worker_loop():195    while True:196        # item = (req: DiarizeRequest, future: Future)197        req, future = await queue.get()198        try:199            async with semaphore:200                # executa processo de diarização em thread (CPU-bound + chamadas CUDA internas)201                loop = asyncio.get_event_loop()202                result = await loop.run_in_executor(None, _process_request, req)203                future.set_result(result)204        except Exception as e:205            future.set_exception(e)206        finally:207            queue.task_done()208 209def _process_request(req: DiarizeRequest) -> DiarizeResponse:210    t0 = time.time()211 212    # 1) decode base64 -> bytes e converter para WAV 16k mono com ffmpeg (rápido)213    try:214        audio_bytes = base64.b64decode(req.audio_base64)215    except Exception:216        raise HTTPException(status_code=400, detail="audio_base64 inválido.")217    wav16 = ffmpeg_bytes_to_wav16k(audio_bytes)218    t_pre = time.time()219 220    # 2) diarização221    segs = run_diarization_on_wav_bytes(wav16)222    # fusões (parâmetros do request ou env defaults)223    gap_tol = req.gap_tol if req.gap_tol is not None else GAP_TOL224    min_dur = req.min_dur if req.min_dur is not None else MIN_DUR225    max_block = req.max_block_sec if req.max_block_sec is not None else MAX_BLOCK226 227    segs = merge_adjacent_same_speaker(segs, gap_tol=gap_tol, min_dur=min_dur)228    segs = fuse_by_duration(segs, max_block_sec=max_block)229    t_diar = time.time()230 231    # 3) cortar e serializar232    items, nspk = slice_segments_to_wav_base64(wav16, segs)233    t_slice = time.time()234 235    return DiarizeResponse(236        segments=sorted(items, key=lambda r: r.start),237        num_speakers=nspk,238        preprocess_sec=round(t_pre - t0, 3),239        diarization_sec=round(t_diar - t_pre, 3),240        slicing_sec=round(t_slice - t_diar, 3),241        total_sec=round(t_slice - t0, 3),242    )243 244# ---------------------------245# Startup: carrega pipeline246# ---------------------------247@app.on_event("startup")248async def startup_worker():249    global PIPELINE250    if not HF_TOKEN:251        # loga e aborta de forma amigável:252        import sys253        print("ERRO: variável HF_TOKEN não definida. Configure em Settings → Repository secrets.", flush=True)254        sys.exit(1)255    print("Start API Diarization")256    # carrega pipeline uma única vez por processo257    PIPELINE = Pipeline.from_pretrained("pyannote/speaker-diarization-3.1", use_auth_token=HF_TOKEN)258 259    # ajustar batch sizes direto nos componentes260    if hasattr(PIPELINE, "segmentation") and hasattr(PIPELINE.segmentation, "batch_size"):261        PIPELINE.segmentation.batch_size = SEG_BS262    if hasattr(PIPELINE, "embedding") and hasattr(PIPELINE.embedding, "batch_size"):263        PIPELINE.embedding.batch_size = EMB_BS264 265    # device266    if torch.cuda.is_available():267        PIPELINE.to(torch.device("cuda"))268        torch.backends.cudnn.benchmark = True269        print(">> usando GPU:", torch.cuda.get_device_name(0))270    else:271        PIPELINE.to(torch.device("cpu"))272        print(">> usando CPU")273 274    # inicia worker275    asyncio.create_task(worker_loop())276 277# ---------------------------278# Endpoints279# ---------------------------280 281@app.post("/diarize/file",282          response_model=DiarizeResponse,283          summary="Diarização via upload de arquivo",284          description="Recebe um arquivo de áudio (multipart/form-data) e retorna os segmentos diarizados.")285async def diarize_file(file: UploadFile = File(...)):286    try:287        audio_bytes = await file.read()288        audio_b64 = base64.b64encode(audio_bytes).decode("utf-8")289        req = DiarizeRequest(audio_base64=audio_b64)290    except Exception:291        raise HTTPException(status_code=400, detail="Erro ao ler arquivo enviado.")292 293    future = asyncio.get_event_loop().create_future()294    await queue.put((req, future))295    result = await future296    return result297 298 299@app.post("/diarize",300          response_model=DiarizeResponse,301          summary="Diarização → chunks WAV base64 por segmento",302          description="Recebe áudio (base64), retorna lista de segmentos com speaker/timestamps e WAV base64.")303async def diarize(req: DiarizeRequest) -> DiarizeResponse:304    future = asyncio.get_event_loop().create_future()305    await queue.put((req, future))306    result = await future307    return result308 309@app.post("/diarize/streaming",310          summary="Diarização em streaming (concatena os trechos WAV no fluxo)",311          description="Recebe áudio base64 e transmite os trechos WAV (concatenados no stream). Útil para consumo progressivo.")312async def diarize_streaming(req: DiarizeRequest):313    future = asyncio.get_event_loop().create_future()314    await queue.put((req, future))315    resp: DiarizeResponse = await future316 317    # streama os WAVs de cada segmento, um atrás do outro318    async def audio_stream():319        for seg in resp.segments:320            raw = base64.b64decode(seg.audio_base64)321            # envia em pedaços (chunked transfer)322            buf = io.BytesIO(raw)323            while True:324                chunk = buf.read(4096)325                if not chunk: break326                yield chunk327 328    # WAV stream (sem cabeçalho único porque são vários trechos encadeados)329    return StreamingResponse(audio_stream(), media_type="application/octet-stream")330 331@app.get("/", response_class=HTMLResponse, include_in_schema=False)332async def index():333    # você pode colocar um index.html em static/ para testar334    return "<h3>API de Diarização pronta</h3>"335 336@app.get("/favicon.ico", include_in_schema=False)337def favicon():338    return Response(status_code=204)339 340if __name__ == "__main__":341    import uvicorn342    port = int(os.getenv("PORT", "7860"))  # Spaces injeta PORT em runtime343    uvicorn.run(app, host="0.0.0.0", port=port)344