Quellcode durchsuchen

feat: add STT endpoint with faster-whisper base model

Xiaogang Liao vor 4 Wochen
Ursprung
Commit
0b8a5bc3d9

+ 5 - 0
cfc-langgraph/Dockerfile

@@ -14,6 +14,11 @@ COPY app/ app/
 COPY src/ src/
 COPY data/ data/
 
+# 预下载 faster-whisper base 模型(避免启动时慢下载)
+ENV HF_ENDPOINT=https://hf-mirror.com
+RUN python3 -c "from faster_whisper import WhisperModel; WhisperModel('base', device='cpu', compute_type='int8')" 2>/dev/null; \
+    chmod -R u+rwX /root/.cache || true
+
 RUN useradd -m -u 1000 appuser && \
     chown -R appuser:appuser /app && \
     mkdir -p /data/chroma_db && \

+ 48 - 0
cfc-langgraph/app/api/audio.py

@@ -0,0 +1,48 @@
+"""语音转写端点"""
+import os
+import uuid
+import logging
+from fastapi import APIRouter, UploadFile, File, HTTPException
+from pydantic import BaseModel
+from typing import Optional
+from app.audio.transcriber import transcribe_audio
+
+logger = logging.getLogger(__name__)
+router = APIRouter(prefix="/api/v1", tags=["audio"])
+
+MAX_FILE_SIZE = 5 * 1024 * 1024
+UPLOAD_DIR = "/tmp/stt_uploads"
+
+
+class TranscribeResponse(BaseModel):
+    code: int = 200
+    message: str = "ok"
+    data: Optional[dict] = None
+
+
+@router.post("/audio/transcribe", response_model=TranscribeResponse)
+async def transcribe(file: UploadFile = File(...)):
+    if not file.filename or not file.filename.lower().endswith((".mp3", ".wav", ".m4a", ".flac", ".ogg")):
+        raise HTTPException(status_code=400, detail="不支持的音频格式,仅支持 mp3/wav/m4a/flac/ogg")
+    content = await file.read()
+    if len(content) > MAX_FILE_SIZE:
+        raise HTTPException(status_code=413, detail="音频文件过大(最大 5MB)")
+    if len(content) < 1000:
+        raise HTTPException(status_code=400, detail="音频文件为空或过短")
+
+    os.makedirs(UPLOAD_DIR, exist_ok=True)
+    ext = file.filename.rsplit(".", 1)[-1]
+    tmp_path = os.path.join(UPLOAD_DIR, f"{uuid.uuid4().hex}.{ext}")
+    try:
+        with open(tmp_path, "wb") as f:
+            f.write(content)
+        result = transcribe_audio(tmp_path)
+        return TranscribeResponse(data=result)
+    except Exception as e:
+        logger.error("转写失败: %s", e, exc_info=True)
+        return TranscribeResponse(code=500, message="转写失败: " + str(e), data={})
+    finally:
+        try:
+            os.remove(tmp_path)
+        except OSError:
+            pass

+ 0 - 0
cfc-langgraph/app/audio/__init__.py


+ 71 - 0
cfc-langgraph/app/audio/transcriber.py

@@ -0,0 +1,71 @@
+"""faster-whisper 本地 STT 引擎封装"""
+import os
+import logging
+import numpy as np
+import av
+from faster_whisper import WhisperModel
+
+logger = logging.getLogger(__name__)
+
+_model = None
+_loaded = False
+
+
+def init_model(model_size: str = "base", device: str = "cpu", compute_type: str = "int8"):
+    """初始化 faster-whisper 模型(启动时调用)"""
+    global _model, _loaded
+    if _loaded:
+        return
+    os.environ.setdefault("HF_ENDPOINT", "https://hf-mirror.com")
+    logger.info("加载 faster-whisper 模型 %s (device=%s, compute=%s)...", model_size, device, compute_type)
+    _model = WhisperModel(model_size, device=device, compute_type=compute_type)
+    _loaded = True
+    logger.info("faster-whisper 模型加载完成")
+
+
+def transcribe_audio(file_path: str, language: str = "zh") -> dict:
+    """转写音频文件,返回 {text, language, language_probability, duration}"""
+    if not _model:
+        raise RuntimeError("faster-whisper 模型未初始化")
+
+    duration_sec = 0.0
+    try:
+        with av.open(file_path) as container:
+            stream = container.streams.audio[0]
+            audio_chunks = []
+            sample_cnt = 0
+            for frame in container.decode(audio=0):
+                arr = frame.to_ndarray()
+                sample_cnt += frame.samples
+                if arr.ndim > 1 and arr.shape[0] > 1:
+                    arr = np.mean(arr, axis=0)
+                else:
+                    arr = arr.reshape(-1)
+                audio_chunks.append(arr)
+            if not audio_chunks:
+                raise ValueError("音频文件无有效帧")
+            audio = np.concatenate(audio_chunks).astype(np.float32)
+            max_val = np.max(np.abs(audio))
+            if max_val > 0:
+                audio = audio / max_val
+            orig_rate = stream.sample_rate if stream.sample_rate else 44100
+            if orig_rate != 16000:
+                new_len = int(sample_cnt * 16000 / orig_rate)
+                audio = np.interp(
+                    np.linspace(0, sample_cnt - 1, new_len),
+                    np.arange(sample_cnt),
+                    audio
+                ).astype(np.float32)
+            duration_sec = len(audio) / 16000.0
+    except Exception as e:
+        logger.error("音频解码失败 %s: %s", file_path, e)
+        raise
+
+    segments, info = _model.transcribe(audio, language=language, beam_size=5)
+    text = "".join(seg.text for seg in segments).strip()
+    return {
+        "text": text,
+        "language": info.language,
+        "language_probability": round(float(info.language_probability), 4),
+        "duration": round(duration_sec, 2),
+    }

+ 6 - 1
cfc-langgraph/app/main.py

@@ -4,7 +4,7 @@ import asyncio
 import logging
 from fastapi import FastAPI, Request
 from starlette.middleware.base import BaseHTTPMiddleware
-from app.api import health, recommend, chat, analyze, tongue, adapter, report_parse, meal, logs
+from app.api import health, recommend, chat, analyze, tongue, adapter, report_parse, meal, logs, audio
 from app import monitoring
 from src.app import router as questionnaire_router
 from app.middleware import request_log_middleware
@@ -26,6 +26,7 @@ app.include_router(monitoring.router)
 app.include_router(questionnaire_router)
 app.include_router(report_parse.router)
 app.include_router(meal.router)
+app.include_router(audio.router)
 app.include_router(logs.router)
 
 
@@ -37,6 +38,10 @@ async def startup():
     json_logs = os.getenv("JSON_LOGS", "false").lower() == "true"
     setup_logging(level=settings.log_level, json_format=json_logs)
 
+    # 初始化 faster-whisper STT 模型
+    from app.audio.transcriber import init_model
+    init_model()
+
     if os.getenv("LANGCHAIN_TRACING_V2", "").lower() == "true":
         logger.info(
             "LangSmith 已启用: project=%s, api_key=%s...",

+ 2 - 0
cfc-langgraph/pyproject.toml

@@ -17,6 +17,8 @@ dependencies = [
     "langchain-chroma==0.1.4",
     "python-multipart>=0.0.20",
     "prometheus-client>=0.21",
+    "faster-whisper==1.2.1",
+    "av==18.1.0",
 ]
 
 [project.optional-dependencies]

+ 2 - 0
cfc-langgraph/requirements.txt

@@ -10,3 +10,5 @@ python-dotenv>=1.0,<2.0
 httpx>=0.27,<1.0
 pytest>=8.0,<9.0
 PyPDF2>=3.0,<4.0
+faster-whisper==1.2.1
+av==18.1.0