Spaces:
Sleeping
Sleeping
| import os | |
| import io | |
| import traceback | |
| import numpy as np | |
| import soundfile as sf | |
| from fastapi import FastAPI, HTTPException | |
| from fastapi.responses import Response | |
| from pydantic import BaseModel | |
| from supertonic import TTS | |
| import uvicorn | |
| app = FastAPI(title="Supertonic TTS API") | |
| # Модели для валидации запросов | |
| class TTSRequest(BaseModel): | |
| text: str | |
| lang: str = "ru" | |
| voice: str = "M2" | |
| # Глобальная загрузка модели | |
| print("Загрузка модели Supertonic TTS...") | |
| tts = TTS(auto_download=True) | |
| default_style = tts.get_voice_style(voice_name="M2") | |
| print("Модель успешно загружена и готова к работе!") | |
| async def root(): | |
| return { | |
| "status": "ok", | |
| "message": "Supertonic TTS API is running", | |
| "docs": "/docs", | |
| "usage": "POST /api/tts с JSON: {'text': 'ваш текст', 'lang': 'ru', 'voice': 'M2'}" | |
| } | |
| async def synthesize(request: TTSRequest): | |
| try: | |
| # Получаем стиль голоса | |
| if request.voice == "M2": | |
| style = default_style | |
| else: | |
| style = tts.get_voice_style(voice_name=request.voice) | |
| # Синтез | |
| wav, duration = tts.synthesize(request.text, voice_style=style, lang=request.lang) | |
| # Конвертация тензоров в numpy если нужно | |
| if hasattr(wav, 'cpu'): | |
| wav = wav.cpu().numpy() | |
| elif hasattr(wav, 'numpy'): | |
| wav = wav.numpy() | |
| wav = np.asarray(wav, dtype=np.float32) | |
| # Получаем sample rate | |
| sample_rate = getattr(tts, 'sample_rate', 24000) | |
| # Записываем в память | |
| out = io.BytesIO() | |
| sf.write(out, wav, samplerate=sample_rate, format='WAV', subtype='PCM_16') | |
| audio_bytes = out.getvalue() | |
| # Возвращаем аудио | |
| return Response( | |
| content=audio_bytes, | |
| media_type='audio/wav', | |
| headers={ | |
| "Content-Disposition": f"attachment; filename=speech.wav", | |
| "X-Audio-Duration": str(round(duration, 2)) | |
| } | |
| ) | |
| except Exception as e: | |
| traceback.print_exc() | |
| raise HTTPException(status_code=500, detail=f"Ошибка генерации: {str(e)}") | |
| if __name__ == '__main__': | |
| port = int(os.environ.get('PORT', 7860)) | |
| uvicorn.run(app, host='0.0.0.0', port=port) |