Files
speech/app/asr/workers/faster_whisper_worker.py
du5t e3b2a53c54 Self-restart worker process when model loading fails
Discovered vibevoice sitting at 3.2GB resident GPU memory with
loaded_models: [] and no idle-unload ever firing for it again. Root cause:
a load attempt had OOM'd partway through (competing with an unrelated
ollama process on the same GPU), so _MODEL_CACHE was never populated —
the idle-unload loop only clears that cache, so it had nothing to act on,
even though the partially-constructed model had already left memory
allocated. gc.collect()+empty_cache() don't reliably reclaim memory from
an interrupted from_pretrained() call.

Confirmed a plain process restart does fully reclaim it, so each worker
now treats any load failure as fatal: log it and os._exit(1), letting
supervisord's autorestart=true respawn a clean process immediately.
Verified with a bogus model id — worker exits, respawns, and passes the
smoke test right after.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-24 00:22:59 +09:00

349 lines
12 KiB
Python

from __future__ import annotations
import argparse
import asyncio
import gc
import tempfile
import time
from pathlib import Path
from typing import Any, Dict, List, Optional, Union
import uvicorn
from fastapi import FastAPI, File, Form, HTTPException, UploadFile
from faster_whisper import WhisperModel
import sys
sys.path.insert(0, "/app")
from asr.config import (
COMPUTE_TYPE,
CUSTOM_MODEL_DIR,
DEFAULT_MODEL,
DEVICE,
MODEL_CACHE,
PYANNOTE_HF_TOKEN,
ensure_runtime_dirs,
resolve_custom_model_path,
)
import os
# 컨테이너는 983 유저로 실행되는데 HOME(/app)이 root 소유라, pyannote.audio가 쓰는
# numba/matplotlib/torch가 각자 ~/.cache, ~/.config 아래에 쓰려다 실패한다.
os.environ["HOME"] = "/tmp" # 컨테이너 기본 HOME=/app은 983 유저가 쓰기 불가 (setdefault로는 덮어쓰기 안 됨)
os.environ.setdefault("NUMBA_CACHE_DIR", "/tmp/numba_cache")
app = FastAPI(title="ASR Faster-Whisper Worker")
_MODEL_CACHE: Dict[str, WhisperModel] = {}
_DIARIZATION_PIPELINE: Any = None
# Idle-unload / model-switch eviction: only one whisper model stays resident
# at a time (switching sizes frees the old one), and everything (including
# the diarization pipeline) is dropped after IDLE_UNLOAD_SECONDS of no
# requests, so this backend doesn't permanently hog GPU memory shared with
# the other ASR workers.
IDLE_UNLOAD_SECONDS = 120
_last_used: float = 0.0
_active_requests: int = 0
def _free_gpu() -> None:
gc.collect()
try:
import torch
if DEVICE == "cuda" and torch.cuda.is_available():
torch.cuda.empty_cache()
except Exception:
pass
def _unload_whisper_models() -> None:
if not _MODEL_CACHE:
return
print(f"[faster_whisper] unloading {list(_MODEL_CACHE.keys())}", flush=True)
_MODEL_CACHE.clear()
_free_gpu()
def _unload_diarization() -> None:
global _DIARIZATION_PIPELINE
if _DIARIZATION_PIPELINE is None:
return
print("[faster_whisper] unloading diarization pipeline", flush=True)
_DIARIZATION_PIPELINE = None
_free_gpu()
async def _idle_unload_loop() -> None:
while True:
await asyncio.sleep(30)
if _active_requests == 0 and _last_used and (time.monotonic() - _last_used) >= IDLE_UNLOAD_SECONDS:
_unload_whisper_models()
_unload_diarization()
@app.on_event("startup")
async def startup() -> None:
ensure_runtime_dirs()
asyncio.create_task(_idle_unload_loop())
@app.get("/health")
def health() -> Dict[str, Any]:
return {
"status": "ok",
"device": DEVICE,
"compute_type": COMPUTE_TYPE,
"diarization_available": bool(PYANNOTE_HF_TOKEN),
}
@app.post("/transcribe")
async def transcribe(
file: UploadFile = File(...),
model: str = Form(DEFAULT_MODEL),
custom_model_path: Optional[str] = Form(None),
language: Optional[str] = Form(None),
task: str = Form("transcribe"),
beam_size: int = Form(5),
temperature: float = Form(0.0),
word_timestamps: bool = Form(False),
diarize: bool = Form(False),
num_speakers: Optional[int] = Form(None),
min_speakers: Optional[int] = Form(None),
max_speakers: Optional[int] = Form(None),
no_repeat_ngram_size: int = Form(0),
repetition_penalty: float = Form(1.0),
compression_ratio_threshold: float = Form(2.4),
log_prob_threshold: float = Form(-1.0),
no_speech_threshold: float = Form(0.6),
condition_on_previous_text: bool = Form(True),
) -> Dict[str, Any]:
ensure_runtime_dirs()
suffix = Path(file.filename or "upload.bin").suffix or ".bin"
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
tmp.write(await file.read())
tmp_path = Path(tmp.name)
global _active_requests, _last_used
_active_requests += 1
try:
resolved_model = resolve_custom_model_path(custom_model_path) or model
result = _transcribe(
audio_path=tmp_path,
model_name=resolved_model,
language=language or None,
task=task,
beam_size=beam_size,
temperature=temperature,
word_timestamps=word_timestamps,
no_repeat_ngram_size=no_repeat_ngram_size,
repetition_penalty=repetition_penalty,
compression_ratio_threshold=compression_ratio_threshold,
log_prob_threshold=log_prob_threshold,
no_speech_threshold=no_speech_threshold,
condition_on_previous_text=condition_on_previous_text,
)
if diarize:
result = _apply_diarization(
audio_path=tmp_path,
result=result,
num_speakers=num_speakers,
min_speakers=min_speakers,
max_speakers=max_speakers,
)
return result
except HTTPException:
raise
except Exception as exc:
raise HTTPException(status_code=500, detail=f"Worker failure: {exc}") from exc
finally:
tmp_path.unlink(missing_ok=True)
_free_gpu()
_active_requests -= 1
_last_used = time.monotonic()
def _load_model(model_name: str) -> WhisperModel:
if model_name not in _MODEL_CACHE:
if _MODEL_CACHE:
# Only one whisper model resident at a time — switching sizes frees the old one.
_unload_whisper_models()
try:
_MODEL_CACHE[model_name] = WhisperModel(
model_name,
device=DEVICE,
compute_type=COMPUTE_TYPE,
download_root=str(MODEL_CACHE),
)
except Exception as e:
# A load interrupted partway (e.g. OOM) can leave CUDA memory
# fragmented/leaked in ways gc.collect()+empty_cache() don't
# reliably reclaim, and since _MODEL_CACHE never got populated
# the idle-unload loop has nothing to clean up either.
# Restarting the whole process is the only guaranteed way to get
# that memory back — supervisord's autorestart=true respawns it
# immediately.
print(f"[faster_whisper] model load failed, restarting process to reclaim GPU memory: {type(e).__name__}: {e}", flush=True)
os._exit(1)
return _MODEL_CACHE[model_name]
def _transcribe(
audio_path: Path,
model_name: str,
language: Optional[str],
task: str,
beam_size: int,
temperature: float,
word_timestamps: bool,
no_repeat_ngram_size: int,
repetition_penalty: float,
compression_ratio_threshold: float,
log_prob_threshold: float,
no_speech_threshold: float,
condition_on_previous_text: bool,
) -> Dict[str, Any]:
model = _load_model(model_name)
temperatures: Union[float, List[float]] = [0.0, 0.2, 0.4, 0.6, 0.8, 1.0] if temperature == 0.0 else temperature
segments_iter, info = model.transcribe(
str(audio_path),
language=language,
task=task,
beam_size=beam_size,
temperature=temperatures,
word_timestamps=word_timestamps,
no_repeat_ngram_size=no_repeat_ngram_size,
repetition_penalty=repetition_penalty,
compression_ratio_threshold=compression_ratio_threshold,
log_prob_threshold=log_prob_threshold,
no_speech_threshold=no_speech_threshold,
condition_on_previous_text=condition_on_previous_text,
)
segment_list = []
full_text = []
for seg in segments_iter:
seg_dict: Dict[str, Any] = {
"id": seg.id,
"start": round(float(seg.start), 3),
"end": round(float(seg.end), 3),
"text": seg.text,
"avg_logprob": round(float(seg.avg_logprob), 4) if seg.avg_logprob is not None else None,
"compression_ratio": round(float(seg.compression_ratio), 4) if seg.compression_ratio is not None else None,
"no_speech_prob": round(float(seg.no_speech_prob), 4) if seg.no_speech_prob is not None else None,
}
if word_timestamps and getattr(seg, "words", None):
seg_dict["words"] = [
{
"start": round(float(w.start), 3),
"end": round(float(w.end), 3),
"word": w.word,
"probability": round(float(w.probability), 4),
}
for w in seg.words
]
segment_list.append(seg_dict)
full_text.append(seg.text.strip())
duration = None
try:
duration = round(float(info.duration), 3)
except Exception:
pass
return {
"backend": "faster-whisper",
"model": model_name,
"language": getattr(info, "language", language),
"language_probability": round(float(getattr(info, "language_probability", 0) or 0), 4),
"duration": duration,
"text": " ".join(x for x in full_text if x),
"segments": segment_list,
"diarized": False,
}
def _get_diarization_pipeline():
global _DIARIZATION_PIPELINE
if _DIARIZATION_PIPELINE is None:
try:
from pyannote.audio import Pipeline
import torch
except ImportError as exc:
raise RuntimeError("pyannote.audio가 설치되지 않았습니다.") from exc
if not PYANNOTE_HF_TOKEN:
raise RuntimeError("화자 분리를 사용하려면 PYANNOTE_HF_TOKEN 환경변수를 설정하세요.")
_DIARIZATION_PIPELINE = Pipeline.from_pretrained(
"pyannote/speaker-diarization-3.1",
token=PYANNOTE_HF_TOKEN,
)
if DEVICE == "cuda":
import torch
_DIARIZATION_PIPELINE = _DIARIZATION_PIPELINE.to(torch.device("cuda"))
return _DIARIZATION_PIPELINE
def _apply_diarization(
audio_path: Path,
result: Dict[str, Any],
num_speakers: Optional[int],
min_speakers: Optional[int],
max_speakers: Optional[int],
) -> Dict[str, Any]:
pipeline = _get_diarization_pipeline()
kwargs: Dict[str, Any] = {}
if num_speakers is not None:
kwargs["num_speakers"] = num_speakers
else:
if min_speakers is not None:
kwargs["min_speakers"] = min_speakers
if max_speakers is not None:
kwargs["max_speakers"] = max_speakers
# pyannote.audio 4.x: pipeline() returns a DiarizeOutput dataclass instead of
# an Annotation directly; the itertracks-capable Annotation is .speaker_diarization.
diarization = pipeline(str(audio_path), **kwargs).speaker_diarization
segments = result.get("segments", [])
for seg in segments:
seg_start = float(seg.get("start", 0))
seg_end = float(seg.get("end", 0))
speaker_times: Dict[str, float] = {}
for turn, _, speaker in diarization.itertracks(yield_label=True):
overlap_start = max(seg_start, turn.start)
overlap_end = min(seg_end, turn.end)
if overlap_end > overlap_start:
speaker_times[speaker] = speaker_times.get(speaker, 0.0) + (overlap_end - overlap_start)
seg["speaker"] = max(speaker_times, key=speaker_times.get) if speaker_times else "UNKNOWN"
lines = []
current_speaker: Optional[str] = None
current_texts: list = []
for seg in segments:
speaker = seg.get("speaker", "UNKNOWN")
text = seg.get("text", "").strip()
if not text:
continue
if speaker != current_speaker:
if current_texts and current_speaker is not None:
lines.append(f"[{current_speaker}]: {' '.join(current_texts)}")
current_speaker = speaker
current_texts = [text]
else:
current_texts.append(text)
if current_texts and current_speaker is not None:
lines.append(f"[{current_speaker}]: {' '.join(current_texts)}")
result["text"] = "\n".join(lines)
result["diarized"] = True
return result
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--host", default="0.0.0.0")
parser.add_argument("--port", type=int, default=8001)
args = parser.parse_args()
ensure_runtime_dirs()
uvicorn.run(app, host=args.host, port=args.port)