Govendas/diarization
0
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 