Files
llm-in-text/backend/tts_asr.py
T

798 lines
25 KiB
Python

# TTS and ASR API for macOS Silicon with HuggingFace transformers
import asyncio
import base64
import logging
import os
import platform
import time
import traceback
from typing import Optional
from fastapi import APIRouter, HTTPException, Security
from pydantic import BaseModel
import numpy as np
router = APIRouter()
logger = logging.getLogger("tts_asr")
# Environment variables
TTS_ASR_DEVICE = os.environ.get("TTS_ASR_DEVICE", "auto")
TTS_ASR_WARMUP = os.environ.get("TTS_ASR_WARMUP", "true").lower() == "true"
TTS_ASR_WARMUP_TIMEOUT = int(os.environ.get("TTS_ASR_WARMUP_TIMEOUT", "120"))
TTS_ASR_IDLE_TIMEOUT = int(os.environ.get("TTS_ASR_IDLE_TIMEOUT", "0"))
# Warmup constants
TTS_WARMUP_TEXT = "你好,这是一个测试。"
ASR_WARMUP_AUDIO_SECONDS = 0.5
# Global state
_tts_pipeline = None
_asr_pipeline = None
_device = None
_device_tested = False
_tts_last_used = 0.0
_asr_last_used = 0.0
_tts_loading = False
_asr_loading = False
_tts_lock = asyncio.Lock()
_asr_lock = asyncio.Lock()
def _test_device_capability(device_str: str) -> tuple[bool, str]:
"""
测试设备实际可用性
返回: (是否可用, 错误信息)
"""
try:
import torch
if device_str == "cpu":
return True, ""
if device_str == "mps":
if not hasattr(torch.backends, "mps") or not torch.backends.mps.is_available():
return False, "MPS 不可用"
if not torch.backends.mps.is_built():
return False, "MPS 未编译"
test_tensor = torch.randn(2, 2, device="mps")
_ = test_tensor @ test_tensor
del test_tensor
torch.mps.empty_cache()
return True, ""
if device_str.startswith("cuda"):
if not torch.cuda.is_available():
return False, "CUDA 不可用"
torch.cuda.empty_cache()
return True, ""
return False, f"未知设备类型: {device_str}"
except Exception as e:
return False, f"设备测试失败: {str(e)}"
def _get_device() -> str:
"""
获取最佳计算设备,支持环境变量覆盖和降级策略
"""
global _device, _device_tested
if _device is not None and _device_tested:
return _device
import torch
device_preference = []
if TTS_ASR_DEVICE == "cpu":
_device = "cpu"
_device_tested = True
logger.info("[Device] 强制使用 CPU (环境变量)")
return _device
elif TTS_ASR_DEVICE in ("mps", "cuda", "auto"):
if TTS_ASR_DEVICE != "auto":
device_preference = [TTS_ASR_DEVICE, "cpu"]
else:
if platform.system() == "Darwin":
device_preference = ["mps", "cpu"]
else:
device_preference = ["cuda", "cpu"]
else:
device_preference = ["mps", "cuda", "cpu"]
for dev in device_preference:
ok, err = _test_device_capability(dev)
if ok:
_device = dev
_device_tested = True
logger.info("[Device] 使用 %s 加速", dev.upper() if dev != "cpu" else "CPU")
return _device
else:
logger.warning("[Device] %s 不可用: %s", dev.upper() if dev != "cpu" else "CPU", err)
_device = "cpu"
_device_tested = True
logger.info("[Device] 降级使用 CPU")
return _device
def _device_arg() -> str:
device = _get_device()
if device == "cuda":
return "cuda:0"
return device
def _get_torch_dtype():
device = _get_device()
import torch
return torch.float16 if device != "cpu" else torch.float32
def _clear_cuda_cache():
try:
import torch
if _device and _device.startswith("cuda"):
torch.cuda.empty_cache()
except Exception:
pass
def _clear_mps_cache():
try:
import torch
if _device == "mps":
torch.mps.empty_cache()
except Exception:
pass
async def _load_tts_pipeline_with_retry(max_retries: int = 2) -> bool:
"""
加载TTS管道,支持重试和降级
"""
global _tts_pipeline, _tts_loading
async with _tts_lock:
if _tts_pipeline is not None:
return True
if _tts_loading:
return False
_tts_loading = True
try:
import torch
from transformers import pipeline
current_device = _get_device()
for attempt in range(max_retries):
try:
device_to_use = _device_arg()
torch_dtype = _get_torch_dtype()
logger.info("[TTS] 加载 Kokoro-82M 模型 (尝试 %d/%d, 设备: %s)...",
attempt + 1, max_retries, device_to_use)
_tts_pipeline = await asyncio.to_thread(
lambda: pipeline(
"text-to-speech",
model="hexgrad/Kokoro-82M",
trust_remote_code=True,
device=device_to_use,
torch_dtype=torch_dtype,
)
)
logger.info("[TTS] Kokoro-82M 模型加载完成")
return True
except RuntimeError as e:
error_str = str(e)
if "MPS" in error_str or "mps" in error_str:
logger.warning("[TTS] MPS 推理失败,尝试降级到 CPU: %s", error_str)
global _device
_device = "cpu"
_clear_mps_cache()
continue
elif "CUDA" in error_str or "cuda" in error_str:
logger.warning("[TTS] CUDA 推理失败,尝试降级到 CPU: %s", error_str)
_device = "cpu"
_clear_cuda_cache()
continue
else:
raise
except Exception as e:
logger.error("[TTS] 加载失败: %s", str(e))
if attempt == max_retries - 1:
raise
await asyncio.sleep(1)
return _tts_pipeline is not None
finally:
_tts_loading = False
async def _load_asr_pipeline_with_retry(max_retries: int = 2) -> bool:
"""
加载ASR管道,支持重试和降级
"""
global _asr_pipeline, _asr_loading
async with _asr_lock:
if _asr_pipeline is not None:
return True
if _asr_loading:
return False
_asr_loading = True
try:
import torch
from transformers import AutoModelForSpeechSeq2Seq, AutoProcessor, pipeline
for attempt in range(max_retries):
try:
device_to_use = _device_arg()
torch_dtype = _get_torch_dtype()
logger.info("[ASR] 加载 Whisper large-v3-turbo 模型 (尝试 %d/%d, 设备: %s)...",
attempt + 1, max_retries, device_to_use)
model_id = "openai/whisper-large-v3-turbo"
def load_model():
model = AutoModelForSpeechSeq2Seq.from_pretrained(
model_id,
torch_dtype=torch_dtype,
low_cpu_mem_usage=True,
use_safetensors=True,
)
processor = AutoProcessor.from_pretrained(model_id)
return pipeline(
"automatic-speech-recognition",
model=model,
tokenizer=processor.tokenizer,
feature_extractor=processor.feature_extractor,
torch_dtype=torch_dtype,
device=device_to_use,
)
_asr_pipeline = await asyncio.to_thread(load_model)
logger.info("[ASR] Whisper large-v3-turbo 模型加载完成")
return True
except RuntimeError as e:
error_str = str(e)
if "MPS" in error_str or "mps" in error_str:
logger.warning("[ASR] MPS 推理失败,尝试降级到 CPU: %s", error_str)
global _device
_device = "cpu"
_clear_mps_cache()
continue
elif "CUDA" in error_str or "cuda" in error_str:
logger.warning("[ASR] CUDA 推理失败,尝试降级到 CPU: %s", error_str)
_device = "cpu"
_clear_cuda_cache()
continue
else:
raise
except Exception as e:
logger.error("[ASR] 加载失败: %s", str(e))
if attempt == max_retries - 1:
raise
await asyncio.sleep(1)
return _asr_pipeline is not None
finally:
_asr_loading = False
def _get_tts_pipeline():
"""同步获取TTS管道(已弃用,保留兼容性)"""
if _tts_pipeline is not None:
return _tts_pipeline
raise RuntimeError("TTS 管道未加载,请使用 _load_tts_pipeline_with_retry()")
def _get_asr_pipeline():
"""同步获取ASR管道(已弃用,保留兼容性)"""
if _asr_pipeline is not None:
return _asr_pipeline
raise RuntimeError("ASR 管道未加载,请使用 _load_asr_pipeline_with_retry()")
async def _warmup_tts() -> bool:
"""
预热TTS模型,减少首次请求延迟
"""
global _tts_last_used
try:
logger.info("[TTS] 开始预热...")
if not await _load_tts_pipeline_with_retry():
logger.error("[TTS] 预热失败:无法加载管道")
return False
tts = _tts_pipeline
if tts is None:
return False
def warmup_inference():
try:
result = tts(TTS_WARMUP_TEXT, voice="af_bella")
if isinstance(result, dict):
audio = result.get("audio")
if hasattr(audio, "cpu"):
_ = audio.cpu()
return True
except Exception as e:
logger.warning("[TTS] 预热推理失败(可忽略): %s", str(e))
return False
success = await asyncio.to_thread(warmup_inference)
_tts_last_used = time.time()
if success:
logger.info("[TTS] 预热完成")
return success
except Exception as e:
logger.error("[TTS] 预热异常: %s", str(e))
traceback.print_exc()
return False
async def _warmup_asr() -> bool:
"""
预热ASR模型,减少首次请求延迟
"""
global _asr_last_used
try:
logger.info("[ASR] 开始预热...")
if not await _load_asr_pipeline_with_retry():
logger.error("[ASR] 预热失败:无法加载管道")
return False
asr = _asr_pipeline
if asr is None:
return False
silence_samples = int(16000 * ASR_WARMUP_AUDIO_SECONDS)
silence_audio = np.zeros(silence_samples, dtype=np.float32)
def warmup_inference():
try:
result = asr(
silence_audio,
sampling_rate=16000,
generate_kwargs={"language": "zh", "task": "transcribe"},
)
return True
except Exception as e:
logger.warning("[ASR] 预热推理失败(可忽略): %s", str(e))
return False
success = await asyncio.to_thread(warmup_inference)
_asr_last_used = time.time()
if success:
logger.info("[ASR] 预热完成")
return success
except Exception as e:
logger.error("[ASR] 预热异常: %s", str(e))
traceback.print_exc()
return False
async def _warmup_all() -> tuple[bool, bool]:
"""
预热所有模型
返回: (TTS预热结果, ASR预热结果)
"""
logger.info("[Warmup] 开始预热所有模型 (超时: %d秒)", TTS_ASR_WARMUP_TIMEOUT)
try:
tts_task = asyncio.create_task(_warmup_tts())
asr_task = asyncio.create_task(_warmup_asr())
done, pending = await asyncio.wait(
[tts_task, asr_task],
timeout=TTS_ASR_WARMUP_TIMEOUT,
return_when=asyncio.ALL_COMPLETED,
)
for task in pending:
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
tts_result = tts_task.result() if tts_task in done else False
asr_result = asr_task.result() if asr_task in done else False
logger.info("[Warmup] 完成: TTS=%s, ASR=%s", tts_result, asr_result)
return tts_result, asr_result
except Exception as e:
logger.error("[Warmup] 异常: %s", str(e))
traceback.print_exc()
return False, False
def _check_and_unload_idle_models():
"""
检查并卸载空闲超过阈值的模型
"""
if TTS_ASR_IDLE_TIMEOUT <= 0:
return
global _tts_pipeline, _asr_pipeline
current_time = time.time()
if _tts_pipeline is not None:
idle_seconds = current_time - _tts_last_used
if idle_seconds > TTS_ASR_IDLE_TIMEOUT:
logger.info("[TTS] 空闲 %.0f 秒,卸载模型", idle_seconds)
_tts_pipeline = None
_clear_cuda_cache()
_clear_mps_cache()
if _asr_pipeline is not None:
idle_seconds = current_time - _asr_last_used
if idle_seconds > TTS_ASR_IDLE_TIMEOUT:
logger.info("[ASR] 空闲 %.0f 秒,卸载模型", idle_seconds)
_asr_pipeline = None
_clear_cuda_cache()
_clear_mps_cache()
def _save_audio_to_wav(audio_data: bytes, sample_rate: int = 16000) -> str:
import tempfile
import wave
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False, mode="wb") as tmp:
with wave.open(tmp.name, "wb") as wf:
wf.setnchannels(1)
wf.setsampwidth(2)
wf.setframerate(sample_rate)
wf.writeframes(audio_data)
return tmp.name
async def _tts_sync_with_retry(text: str, voice: str = "af_bella", rate: float = 1.0, max_retries: int = 2) -> tuple[bytes, int]:
"""
TTS推理,支持重试和降级
"""
global _tts_last_used
_check_and_unload_idle_models()
if not await _load_tts_pipeline_with_retry():
raise RuntimeError("TTS 模型加载失败")
tts = _tts_pipeline
sample_rate = 24000
for attempt in range(max_retries):
try:
def inference():
result = tts(text, voice=voice)
audio = None
sr = sample_rate
if isinstance(result, dict):
audio = result.get("audio")
sr = int(result.get("sampling_rate", sr))
elif isinstance(result, (list, tuple)) and result:
audio = result[0]
if audio is None:
raise RuntimeError("Kokoro 未返回音频数据")
if hasattr(audio, "cpu"):
audio = audio.cpu().numpy()
if hasattr(audio, "numpy"):
audio = audio.numpy()
return audio, sr
audio, sample_rate = await asyncio.to_thread(inference)
duration_ms = int(len(audio) * 1000 / sample_rate)
if audio.dtype != np.int16:
audio = (audio * 32767).astype(np.int16)
import tempfile
import wave
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
output_path = tmp.name
try:
with wave.open(output_path, "wb") as wf:
wf.setnchannels(1)
wf.setsampwidth(2)
wf.setframerate(sample_rate)
wf.writeframes(audio.tobytes())
with open(output_path, "rb") as f:
audio_bytes = f.read()
_tts_last_used = time.time()
return audio_bytes, duration_ms
finally:
if os.path.exists(output_path):
os.unlink(output_path)
except RuntimeError as e:
error_str = str(e)
if "MPS" in error_str or "mps" in error_str:
logger.warning("[TTS] MPS 推理错误,尝试降级重试 (尝试 %d/%d): %s",
attempt + 1, max_retries, error_str)
global _device
_device = "cpu"
_clear_mps_cache()
if attempt < max_retries - 1:
continue
elif "CUDA" in error_str or "cuda" in error_str:
logger.warning("[TTS] CUDA 推理错误,尝试降级重试 (尝试 %d/%d): %s",
attempt + 1, max_retries, error_str)
_device = "cpu"
_clear_cuda_cache()
if attempt < max_retries - 1:
continue
raise
except Exception as e:
logger.error("[TTS] 推理失败: %s", str(e))
if attempt == max_retries - 1:
raise
await asyncio.sleep(0.5)
raise RuntimeError("TTS 推理失败")
async def _asr_sync_with_retry(audio_data: bytes, language: str = "zh", max_retries: int = 2) -> str:
"""
ASR推理,支持重试和降级
"""
global _asr_last_used
_check_and_unload_idle_models()
if not await _load_asr_pipeline_with_retry():
raise RuntimeError("ASR 模型加载失败")
audio_path = _save_audio_to_wav(audio_data)
try:
import soundfile as sf
audio_array, sample_rate = await asyncio.to_thread(lambda: sf.read(audio_path))
if len(audio_array.shape) > 1:
audio_array = np.mean(audio_array, axis=1)
if sample_rate != 16000:
import librosa
audio_array = await asyncio.to_thread(
lambda: librosa.resample(audio_array, orig_sr=sample_rate, target_sr=16000)
)
sample_rate = 16000
audio_array = audio_array.astype(np.float32)
for attempt in range(max_retries):
try:
def inference():
asr = _asr_pipeline
result = asr(
audio_array,
sampling_rate=sample_rate,
generate_kwargs={"language": language, "task": "transcribe"},
)
if isinstance(result, dict):
return result.get("text", "").strip()
return str(result).strip()
text = await asyncio.to_thread(inference)
_asr_last_used = time.time()
return text
except RuntimeError as e:
error_str = str(e)
if "MPS" in error_str or "mps" in error_str:
logger.warning("[ASR] MPS 推理错误,尝试降级重试 (尝试 %d/%d): %s",
attempt + 1, max_retries, error_str)
global _device
_device = "cpu"
_clear_mps_cache()
if attempt < max_retries - 1:
continue
elif "CUDA" in error_str or "cuda" in error_str:
logger.warning("[ASR] CUDA 推理错误,尝试降级重试 (尝试 %d/%d): %s",
attempt + 1, max_retries, error_str)
_device = "cpu"
_clear_cuda_cache()
if attempt < max_retries - 1:
continue
raise
except Exception as e:
logger.error("[ASR] 推理失败: %s", str(e))
if attempt == max_retries - 1:
raise
await asyncio.sleep(0.5)
raise RuntimeError("ASR 推理失败")
finally:
if os.path.exists(audio_path):
os.unlink(audio_path)
# Legacy sync wrappers (for compatibility)
def _tts_sync(text: str, voice: str = "af_bella", rate: float = 1.0) -> tuple[bytes, int]:
raise RuntimeError("请使用 _tts_sync_with_retry()")
def _asr_sync(audio_data: bytes, language: str = "zh") -> str:
raise RuntimeError("请使用 _asr_sync_with_retry()")
async def _text_to_speech(text: str, voice: str = "af_bella", rate: float = 1.0) -> tuple[bytes, int]:
return await _tts_sync_with_retry(text, voice, rate)
async def _speech_to_text(audio_data: bytes, language: str = "zh") -> str:
return await _asr_sync_with_retry(audio_data, language)
# Request/Response models
class TTSRequest(BaseModel):
text: str
voice: str = "af_bella"
rate: float = 1.0
format: str = "wav"
class TTSResponse(BaseModel):
audio_base64: str
format: str
duration_ms: int
class ASRRequest(BaseModel):
audio_base64: str
language: str = "zh-CN"
class ASRResponse(BaseModel):
text: str
language: str
class ModelStatus(BaseModel):
tts_loaded: bool
asr_loaded: bool
device: str
tts_last_used: Optional[float] = None
asr_last_used: Optional[float] = None
def get_api_key(api_key: str):
import main
API_KEY = main.API_KEY
if api_key != API_KEY:
raise HTTPException(status_code=403, detail="API Key 无效")
return api_key
@router.get("/status", response_model=ModelStatus)
async def get_status(api_key: str = Security(get_api_key)):
"""
获取模型状态
"""
current_time = time.time()
return ModelStatus(
tts_loaded=_tts_pipeline is not None,
asr_loaded=_asr_pipeline is not None,
device=_get_device(),
tts_last_used=_tts_last_used if _tts_last_used > 0 else None,
asr_last_used=_asr_last_used if _asr_last_used > 0 else None,
)
@router.post("/warmup")
async def warmup_models(api_key: str = Security(get_api_key)):
"""
手动触发模型预热
"""
tts_result, asr_result = await _warmup_all()
return {
"tts_warmup": tts_result,
"asr_warmup": asr_result,
"device": _get_device(),
}
@router.post("/tts", response_model=TTSResponse)
async def text_to_speech(req: TTSRequest, api_key: str = Security(get_api_key)):
request_id = str(hash(req.text))[:8]
try:
logger.info("[TTS][%s] text_chars=%d voice=%s format=%s",
request_id, len(req.text), req.voice, req.format)
audio_data, duration_ms = await _text_to_speech(req.text, req.voice, req.rate)
if req.format.lower() == "mp3":
import subprocess
import tempfile
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_in:
tmp_in.write(audio_data)
input_path = tmp_in.name
with tempfile.NamedTemporaryFile(suffix=".mp3", delete=False) as tmp_out:
output_path = tmp_out.name
try:
cmd = ["ffmpeg", "-i", input_path, "-acodec", "libmp3lame", "-ab", "128k", output_path]
result = await asyncio.to_thread(
lambda: subprocess.run(cmd, capture_output=True, text=True, timeout=30)
)
if result.returncode != 0:
raise RuntimeError(f"MP3 转换失败: {result.stderr}")
with open(output_path, "rb") as f:
audio_data = f.read()
finally:
for path in [input_path, output_path]:
if os.path.exists(path):
os.unlink(path)
logger.info("[TTS][%s] success duration_ms=%d", request_id, duration_ms)
return TTSResponse(
audio_base64=base64.b64encode(audio_data).decode(),
format=req.format,
duration_ms=duration_ms,
)
except Exception as e:
logger.exception("[TTS] failed: %s", e)
raise HTTPException(status_code=500, detail=str(e))
@router.post("/asr", response_model=ASRResponse)
async def speech_to_text(req: ASRRequest, api_key: str = Security(get_api_key)):
request_id = str(hash(req.audio_base64))[:8]
try:
logger.info("[ASR][%s] audio_base64_chars=%d language=%s",
request_id, len(req.audio_base64), req.language)
audio_data = base64.b64decode(req.audio_base64)
text = await _speech_to_text(audio_data, req.language[:2])
logger.info("[ASR][%s] success text_chars=%d", request_id, len(text))
return ASRResponse(text=text, language=req.language)
except Exception as e:
logger.exception("[ASR] failed: %s", e)
raise HTTPException(status_code=500, detail=str(e))
def register_tts_asr_routes(app):
"""
注册TTS/ASR路由并可选执行预热
"""
app.include_router(router, prefix="/v1/tts-asr")
if TTS_ASR_WARMUP:
@app.on_event("startup")
async def warmup_on_startup():
logger.info("[Startup] 开始后台预热...")
asyncio.create_task(_warmup_all())