Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions backend/fastrtc/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
from .speech_to_text import MoonshineSTT, get_stt_model
from .stream import Stream, UIArgs
from .text_to_speech import (
CambTTSOptions,
CartesiaTTSOptions,
KokoroTTSOptions,
get_tts_model,
Expand Down Expand Up @@ -92,6 +93,7 @@
"VideoStreamHandler",
"CloseStream",
"get_current_context",
"CambTTSOptions",
"CartesiaTTSOptions",
"WebRTCData",
]
3 changes: 2 additions & 1 deletion backend/fastrtc/text_to_speech/__init__.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,8 @@
from .tts import (
CambTTSOptions,
CartesiaTTSOptions,
KokoroTTSOptions,
get_tts_model,
)

__all__ = ["get_tts_model", "KokoroTTSOptions", "CartesiaTTSOptions"]
__all__ = ["get_tts_model", "KokoroTTSOptions", "CartesiaTTSOptions", "CambTTSOptions"]
84 changes: 83 additions & 1 deletion backend/fastrtc/text_to_speech/tts.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ class KokoroTTSOptions(TTSOptions):

@lru_cache
def get_tts_model(
model: Literal["kokoro", "cartesia"] = "kokoro", **kwargs
model: Literal["kokoro", "cartesia", "camb"] = "kokoro", **kwargs
) -> TTSModel:
if model == "kokoro":
m = KokoroTTSModel()
Expand All @@ -52,6 +52,9 @@ def get_tts_model(
elif model == "cartesia":
m = CartesiaTTSModel(api_key=kwargs.get("cartesia_api_key", ""))
return m
elif model == "camb":
m = CambTTSModel(api_key=kwargs.get("camb_api_key", ""))
return m
else:
raise ValueError(f"Invalid model: {model}")

Expand Down Expand Up @@ -162,6 +165,85 @@ class CartesiaTTSOptions(TTSOptions):
sample_rate: int = 22_050


@dataclass
class CambTTSOptions(TTSOptions):
voice_id: int = 2681
language: str = "en-us"
model: str = "mars-flash"
speed: float = 1.0
output_format: str = "pcm_s16le"
user_instructions: str | None = None


class CambTTSModel(TTSModel):
def __init__(self, api_key: str):
if importlib.util.find_spec("camb") is None:
raise RuntimeError(
"camb is not installed. Please install it using 'pip install camb'."
)
self._api_key = api_key

def _build_tts_kwargs(self, text: str, options: CambTTSOptions):
kwargs = {
"text": text,
"language": options.language,
"voice_id": options.voice_id,
"speech_model": options.model,
"output_configuration": {"format": options.output_format},
"voice_settings": {"speed": options.speed},
}
if options.model == "mars-instruct" and options.user_instructions:
kwargs["user_instructions"] = options.user_instructions
return kwargs

async def stream_tts(
self, text: str, options: CambTTSOptions | None = None
) -> AsyncGenerator[tuple[int, NDArray[np.int16]], None]:
from camb.client import AsyncCambAI

options = options or CambTTSOptions()
client = AsyncCambAI(api_key=self._api_key)

sentences = re.split(r"(?<=[.!?])\s+", text.strip())

for sentence in sentences:
if not sentence.strip():
continue
async for output in async_aggregate_bytes_to_16bit(
client.text_to_speech.tts(**self._build_tts_kwargs(sentence, options))
):
yield 24000, output.flatten()

def stream_tts_sync(
self, text: str, options: CambTTSOptions | None = None
) -> Generator[tuple[int, NDArray[np.int16]], None, None]:
loop = asyncio.new_event_loop()

iterator = self.stream_tts(text, options).__aiter__()
while True:
try:
yield loop.run_until_complete(iterator.__anext__())
except StopAsyncIteration:
break

def tts(
self, text: str, options: CambTTSOptions | None = None
) -> tuple[int, NDArray[np.int16]]:
loop = asyncio.new_event_loop()
buffer = np.array([], dtype=np.int16)

options = options or CambTTSOptions()

iterator = self.stream_tts(text, options).__aiter__()
while True:
try:
_, chunk = loop.run_until_complete(iterator.__anext__())
buffer = np.concatenate([buffer, chunk])
except StopAsyncIteration:
break
return 24000, buffer


class CartesiaTTSModel(TTSModel):
def __init__(self, api_key: str):
if importlib.util.find_spec("cartesia") is None:
Expand Down
114 changes: 114 additions & 0 deletions demo/camb_voice_agent/app.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,114 @@
import json
import os
import time
from pathlib import Path

import gradio as gr
import numpy as np
from dotenv import load_dotenv
from fastapi import FastAPI
from fastapi.responses import HTMLResponse, StreamingResponse
from fastrtc import (
AdditionalOutputs,
CambTTSOptions,
ReplyOnPause,
Stream,
)
from fastrtc.text_to_speech.tts import CambTTSModel
from fastrtc.utils import audio_to_bytes
from openai import OpenAI
from pydantic import BaseModel

load_dotenv()

openai_client = OpenAI()
tts_model = CambTTSModel(api_key=os.environ["CAMB_API_KEY"])
tts_options = CambTTSOptions(voice_id=int(os.environ.get("CAMB_VOICE_ID", "156549")))

curr_dir = Path(__file__).parent


def response(
audio: tuple[int, np.ndarray],
chatbot: list[dict] | None = None,
):
chatbot = chatbot or []
messages = [{"role": d["role"], "content": d["content"]} for d in chatbot]

prompt = openai_client.audio.transcriptions.create(
file=("audio-file.mp3", audio_to_bytes(audio)),
model="whisper-1",
).text
chatbot.append({"role": "user", "content": prompt})
yield AdditionalOutputs(chatbot)

messages.append({"role": "user", "content": prompt})
llm_response = openai_client.chat.completions.create(
model="gpt-4o-mini",
max_tokens=512,
messages=messages,
)
response_text = llm_response.choices[0].message.content or ""
chatbot.append({"role": "assistant", "content": response_text})

start = time.time()
print("starting tts", start)
for i, chunk in enumerate(tts_model.stream_tts_sync(response_text, tts_options)):
print("chunk", i, time.time() - start)
yield chunk
print("finished tts", time.time() - start)
yield AdditionalOutputs(chatbot)


chatbot = gr.Chatbot(type="messages")
stream = Stream(
modality="audio",
mode="send-receive",
handler=ReplyOnPause(response),
additional_outputs_handler=lambda a, b: b,
additional_inputs=[chatbot],
additional_outputs=[chatbot],
)


class Message(BaseModel):
role: str
content: str


class InputData(BaseModel):
webrtc_id: str
chatbot: list[Message]


app = FastAPI()
stream.mount(app)


@app.get("/")
async def _():
html_content = (curr_dir / "index.html").read_text()
html_content = html_content.replace("__RTC_CONFIGURATION__", json.dumps(None))
return HTMLResponse(content=html_content, status_code=200)


@app.post("/input_hook")
async def _(body: InputData):
stream.set_input(body.webrtc_id, body.model_dump()["chatbot"])
return {"status": "ok"}


@app.get("/outputs")
def _(webrtc_id: str):
async def output_stream():
async for output in stream.output_stream(webrtc_id):
chatbot = output.args[0]
yield f"event: output\ndata: {json.dumps(chatbot[-1])}\n\n"

return StreamingResponse(output_stream(), media_type="text/event-stream")


if __name__ == "__main__":
import uvicorn

uvicorn.run(app, host="0.0.0.0", port=7860)
Loading