1+ #!/usr/bin/env python3
2+ """
3+ Sesame CSM TTS Server
4+ Provides text-to-speech using Sesame's Conversational Speech Model
5+ """
6+
7+ import os
8+ import logging
9+ import io
10+ import base64
11+ from typing import Optional
12+ from datetime import datetime
13+
14+ import torch
15+ import torchaudio
16+ from fastapi import FastAPI , HTTPException , Request
17+ from fastapi .middleware .cors import CORSMiddleware
18+ from pydantic import BaseModel
19+ from transformers import AutoProcessor , CsmForConditionalGeneration
20+ import uvicorn
21+
22+ # Configure logging
23+ logging .basicConfig (level = logging .INFO )
24+ logger = logging .getLogger (__name__ )
25+
26+ # Create FastAPI app
27+ app = FastAPI (
28+ title = "Sesame CSM TTS Server" ,
29+ description = "Text-to-speech using Sesame's Conversational Speech Model" ,
30+ version = "1.0.0"
31+ )
32+
33+ # Add CORS middleware
34+ app .add_middleware (
35+ CORSMiddleware ,
36+ allow_origins = ["*" ],
37+ allow_credentials = True ,
38+ allow_methods = ["GET" , "POST" ],
39+ allow_headers = ["*" ],
40+ )
41+
42+ class TTSRequest (BaseModel ):
43+ text : str
44+ speaker : Optional [int ] = 0
45+ max_audio_length_ms : Optional [int ] = 10000
46+ return_format : Optional [str ] = "wav" # wav, mp3, base64
47+
48+ class TTSResponse (BaseModel ):
49+ audio_data : str # base64 encoded audio
50+ sample_rate : int
51+ format : str
52+ duration_ms : float
53+
54+ class TTSServer :
55+ def __init__ (self ):
56+ self .model = None
57+ self .processor = None
58+ self .device = self ._detect_device ()
59+ self .sample_rate = 24000 # CSM default sample rate
60+
61+ logger .info (f"Using device: { self .device } " )
62+
63+ def _detect_device (self ):
64+ """Detect the best available device."""
65+ if torch .cuda .is_available ():
66+ return "cuda"
67+ elif torch .backends .mps .is_available ():
68+ return "mps"
69+ else :
70+ return "cpu"
71+
72+ def load_model (self ):
73+ """Load the Sesame CSM model."""
74+ try :
75+ model_id = "sesame/csm-1b"
76+ logger .info (f"Loading Sesame CSM model: { model_id } " )
77+
78+ # Load processor and model
79+ self .processor = AutoProcessor .from_pretrained (model_id )
80+ self .model = CsmForConditionalGeneration .from_pretrained (
81+ model_id ,
82+ device_map = self .device ,
83+ torch_dtype = torch .float16 if self .device == "cuda" else torch .float32
84+ )
85+
86+ logger .info ("✓ Sesame CSM model loaded successfully" )
87+ return True
88+
89+ except Exception as e :
90+ logger .error (f"Failed to load model: { e } " )
91+ return False
92+
93+ def generate_audio (self , text : str , speaker : int = 0 , max_audio_length_ms : int = 10000 ):
94+ """Generate audio from text using CSM."""
95+ try :
96+ if not self .model or not self .processor :
97+ raise ValueError ("Model not loaded" )
98+
99+ logger .info (f"Generating audio for text: '{ text [:50 ]} ...'" )
100+
101+ # Prepare inputs
102+ inputs = self .processor (
103+ text = text ,
104+ speaker_id = speaker ,
105+ return_tensors = "pt"
106+ ).to (self .device )
107+
108+ # Generate audio
109+ with torch .no_grad ():
110+ audio_codes = self .model .generate (
111+ ** inputs ,
112+ max_new_tokens = max_audio_length_ms // 25 , # Rough estimate
113+ do_sample = True ,
114+ temperature = 0.7
115+ )
116+
117+ # Decode audio
118+ audio_array = self .processor .decode (audio_codes [0 ])
119+
120+ # Ensure audio is on CPU and in correct format
121+ if isinstance (audio_array , torch .Tensor ):
122+ audio_array = audio_array .cpu ().float ()
123+
124+ logger .info (f"✓ Generated audio: { audio_array .shape } samples at { self .sample_rate } Hz" )
125+ return audio_array
126+
127+ except Exception as e :
128+ logger .error (f"Audio generation failed: { e } " )
129+ raise
130+
131+ def audio_to_base64 (self , audio_tensor , format = "wav" ):
132+ """Convert audio tensor to base64 encoded string."""
133+ try :
134+ # Create a bytes buffer
135+ buffer = io .BytesIO ()
136+
137+ # Save audio to buffer
138+ torchaudio .save (
139+ buffer ,
140+ audio_tensor .unsqueeze (0 ) if audio_tensor .dim () == 1 else audio_tensor ,
141+ self .sample_rate ,
142+ format = format
143+ )
144+
145+ # Get bytes and encode to base64
146+ buffer .seek (0 )
147+ audio_bytes = buffer .getvalue ()
148+ audio_base64 = base64 .b64encode (audio_bytes ).decode ('utf-8' )
149+
150+ return audio_base64
151+
152+ except Exception as e :
153+ logger .error (f"Audio encoding failed: { e } " )
154+ raise
155+
156+ # Global TTS server instance
157+ tts_server = TTSServer ()
158+
159+ @app .on_event ("startup" )
160+ async def startup_event ():
161+ """Load model on startup."""
162+ logger .info ("🚀 Starting Sesame CSM TTS Server" )
163+
164+ if not tts_server .load_model ():
165+ logger .error ("❌ Failed to load TTS model" )
166+ raise RuntimeError ("Model loading failed" )
167+
168+ logger .info ("✅ TTS Server ready!" )
169+
170+ @app .post ("/v1/speak" , response_model = TTSResponse )
171+ async def speak (request : TTSRequest , http_request : Request ):
172+ """Generate speech from text (Deepgram-compatible endpoint)."""
173+
174+ client_ip = http_request .client .host if http_request .client else "unknown"
175+ logger .info (f"TTS request from { client_ip } : '{ request .text [:50 ]} ...'" )
176+
177+ try :
178+ # Generate audio
179+ audio_tensor = tts_server .generate_audio (
180+ text = request .text ,
181+ speaker = request .speaker ,
182+ max_audio_length_ms = request .max_audio_length_ms
183+ )
184+
185+ # Convert to base64
186+ audio_base64 = tts_server .audio_to_base64 (audio_tensor , request .return_format )
187+
188+ # Calculate duration
189+ duration_ms = len (audio_tensor ) / tts_server .sample_rate * 1000
190+
191+ return TTSResponse (
192+ audio_data = audio_base64 ,
193+ sample_rate = tts_server .sample_rate ,
194+ format = request .return_format ,
195+ duration_ms = duration_ms
196+ )
197+
198+ except Exception as e :
199+ logger .error (f"TTS generation failed: { e } " )
200+ raise HTTPException (status_code = 500 , detail = str (e ))
201+
202+ @app .post ("/tts" )
203+ async def tts_simple (request : TTSRequest ):
204+ """Simple TTS endpoint."""
205+ return await speak (request , None )
206+
207+ @app .get ("/health" )
208+ async def health_check ():
209+ """Health check endpoint."""
210+ model_loaded = tts_server .model is not None
211+ return {
212+ "status" : "healthy" if model_loaded else "unhealthy" ,
213+ "service" : "sesame-csm-tts" ,
214+ "timestamp" : datetime .utcnow ().isoformat () + "Z" ,
215+ "model_loaded" : model_loaded ,
216+ "device" : tts_server .device ,
217+ "sample_rate" : tts_server .sample_rate
218+ }
219+
220+ @app .get ("/" )
221+ async def root ():
222+ """Root endpoint with API information."""
223+ return {
224+ "service" : "Sesame CSM TTS Server" ,
225+ "version" : "1.0.0" ,
226+ "description" : "Text-to-speech using Sesame's Conversational Speech Model" ,
227+ "endpoints" : {
228+ "POST /v1/speak" : "Generate speech (Deepgram-compatible)" ,
229+ "POST /tts" : "Simple TTS generation" ,
230+ "GET /health" : "Health check" ,
231+ "GET /" : "This information"
232+ },
233+ "docs" : "/docs"
234+ }
235+
236+ if __name__ == "__main__" :
237+ port = int (os .getenv ("PORT" , 8001 ))
238+ logger .info (f"Starting Sesame CSM TTS Server on port { port } " )
239+ uvicorn .run (app , host = "0.0.0.0" , port = port )
0 commit comments