/
mishzorikhin
/
VoiceDecodeAPI
Обзор
Документация
Войти
/
mishzorikhin
/
VoiceDecodeAPI
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
app.py
243 строки
10 KB
mish
first commit
04 янв 2025, 18:25
04 янв 2025, 18:25
4327bcd
Код
Авторство
О чём код?
import asyncio import os import logging from contextlib import asynccontextmanager from tempfile import NamedTemporaryFile from typing import Optional, Literal, Union import torch import whisper from fastapi import FastAPI, UploadFile, HTTPException, Query, Request from fastapi.encoders import jsonable_encoder from starlette.responses import JSONResponse from models import JSONResponseModel, TextResponse from services.cache_service import init_db, clear_cache from services.whisper_service import transcribe_audio # Настройка логирования logging.basicConfig( level=logging.INFO, format="%(asctime)s - %(levelname)s - %(name)s - %(message)s" ) logger = logging.getLogger(__name__) INACTIVITY_TIMEOUT = 300 model_name = os.environ.get("MODEL_DEFAULT", "small") logger.info(f"Используемая модель: {model_name}") CACHE_DIR = os.path.expanduser("~/.cache/whisper/") os.makedirs(CACHE_DIR, exist_ok=True) logger.info(f"Каталог кэша: {CACHE_DIR}") # Полный путь к модели model_path = os.path.join(CACHE_DIR, f"{model_name}.pt") @asynccontextmanager async def lifespan(app: FastAPI): """ Управление жизненным циклом приложения FastAPI. Загрузка модели при старте и настройка механизма выгрузки. """ model = None device = "cpu" inactivity_timer = None # Таймер для отслеживания бездействия async def unload_model(): """Функция для выгрузки модели при бездействии.""" nonlocal model if model is not None: logger.info("Выгрузка модели из памяти из-за бездействия...") model = None app.state.model = None torch.cuda.empty_cache() logger.info("Модель успешно выгружена.") async def reset_inactivity_timer(): """Перезапуск таймера бездействия.""" nonlocal inactivity_timer if inactivity_timer: inactivity_timer.cancel() inactivity_timer = asyncio.create_task( asyncio.sleep(INACTIVITY_TIMEOUT) ) inactivity_timer.add_done_callback(lambda _: asyncio.create_task(unload_model())) async def load_model_if_needed(): """Загрузка модели, если она не в памяти.""" nonlocal model, device if model is None: logger.info("Загрузка модели...") model = whisper.load_model(model_name) device = "cuda" if torch.cuda.is_available() else "cpu" model.to(device) app.state.model = model app.state.device = device logger.info(f"Модель '{model_name}' загружена на устройство '{device}'.") # Настройка приложения app.state.model = None app.state.device = "cpu" app.state.reset_inactivity_timer = reset_inactivity_timer app.state.load_model_if_needed = load_model_if_needed yield # Выгрузка модели при завершении приложения if inactivity_timer: inactivity_timer.cancel() await unload_model() app = FastAPI( title="Whisper Transcription API", description="API для транскрипции аудиофайлов с использованием модели Whisper.", version="1.0.0", lifespan=lifespan, ) # Инициализация базы данных кэша try: init_db() logger.info("База данных кэша успешно инициализирована.") except Exception as e: logger.error(f"Ошибка при инициализации базы данных кэша: {e}") raise @app.middleware("http") async def reset_timer_on_request(request: Request, call_next): """Перезапуск таймера бездействия при каждом запросе.""" await app.state.reset_inactivity_timer() return await call_next(request) @app.get("/health", summary="Проверка состояния системы", description="Возвращает статус работы сервера.") async def health_check(): model_loaded = app.state.model is not None return {"status": "ok", "model_loaded": model_loaded, "device": app.state.device if model_loaded else None} @app.post( "/transcribe", summary="Транскрибировать аудиофайл", description="Принимает аудиофайл и возвращает его транскрипцию.", response_model=Union[TextResponse, JSONResponseModel], responses={ 400: {"description": "Неправильный тип файла."}, 500: {"description": "Внутренняя ошибка сервера."} } ) async def transcribe_endpoint( file: UploadFile, desc: Optional[str] = Query(None, description="Начальный текст для подсказки модели (опционально)."), response_format: Literal["text", "json"] = Query("text", description='Формат ответа ("text" или "json").'), temperature: Optional[float] = Query(0.0, description="Температура (опционально)."), use_cache: Optional[bool] = Query(True, description="Использовать кэш для транскрипции (опционально).") ) -> JSONResponse: """ Обрабатывает загруженный аудиофайл и возвращает его транскрипцию. :param use_cache: Использовать кэш для транскрипции (опционально). :param temperature: Температура :param file: Загружаемый аудиофайл. :param desc: Начальный текст для подсказки модели (опционально). :param response_format: Формат ответа ("text" или "json"). :return: Результат транскрипции в указанном формате. """ logger.info("Получен запрос на транскрипцию аудиофайла.") await app.state.load_model_if_needed() # Проверка типа контента файла if not file.content_type.startswith("audio/"): logger.warning(f"Неподдерживаемый тип файла: {file.content_type}") raise HTTPException( status_code=400, detail="Загруженный файл не является аудиофайлом." ) # Создание временного файла для сохранения загруженного аудиофайла try: with NamedTemporaryFile(delete=False, suffix=".mp3") as temp_file: content = await file.read() temp_file.write(content) temp_file_path = temp_file.name logger.info(f"Временный файл сохранён: {temp_file_path}") except Exception as e: logger.error(f"Ошибка при сохранении временного файла: {e}") raise HTTPException( status_code=500, detail="Не удалось сохранить загруженный файл." ) # Обработка транскрипции try: logger.info("Начало транскрипции аудиофайла.") transcription_result = await transcribe_audio( temp_file_path=temp_file_path, model=app.state.model, device=app.state.device, language=None, initial_prompt=desc, temperature=temperature, use_cache=use_cache, response_format=response_format ) logger.info("Транскрипция успешно завершена.") except Exception as e: logger.error(f"Ошибка при транскрипции аудиофайла: {e}") raise HTTPException( status_code=500, detail="Ошибка при транскрипции аудиофайла." ) finally: # Удаление временного файла try: os.remove(temp_file_path) logger.info(f"Временный файл удалён: {temp_file_path}") except Exception as e: logger.warning(f"Не удалось удалить временный файл: {temp_file_path}. Ошибка: {e}") return JSONResponse(content=jsonable_encoder(transcription_result)) @app.delete("/cache/clear", summary="Очистить кэш", description="Удаляет все записи из базы данных кэша.") async def clear_cache_endpoint(): try: clear_cache() return {"message": "The cache has been cleared successfully."} except Exception as e: logger.error(f"Ошибка при очистке кэша: {e}") raise HTTPException(status_code=500, detail="Failed to clear cache.") @app.get("/model/status", summary="Получить статус модели", description="Возвращает статус модели и устройства выполнения.") async def get_model_status(): model = app.state.model device = app.state.device return { "model_name": model_name, "device": device, "cuda_available": torch.cuda.is_available() } @app.get("/models", summary="Список скачанных моделей", description="Возвращает список всех локально доступных моделей.") async def list_downloaded_models(): try: if not os.path.exists(CACHE_DIR): raise HTTPException(status_code=404, detail="Директория с кэшем моделей не найдена.") model_files = [file for file in os.listdir(CACHE_DIR) if file.endswith(".pt")] if not model_files: return {"models": []} return {"models": model_files} except Exception as e: raise HTTPException(status_code=500, detail=f"Error getting list of models: {e}")