583 / server.py
Pranavz's picture
Add standalone HTTP and UI server
23ab89a verified
Raw
History Blame Contribute Delete
7.04 kB
#!/usr/bin/env python3
"""Standalone TinyAya HTTP server with a minimal browser UI."""
import argparse
import asyncio
import io
import re
import wave
from pathlib import Path
import torch
from fastapi import FastAPI, HTTPException
from fastapi.responses import HTMLResponse, Response
from huggingface_hub import snapshot_download
from pydantic import BaseModel, Field
from transformers import AutoFeatureExtractor, AutoModelForCausalLM, AutoTokenizer, MimiModel
import uvicorn
AUDIO_RE = re.compile(r"^<(\d+)_(\d+)>$")
SPEAKERS = ("Ira", "Aisha", "Siya", "Zoya", "Silver")
class SpeechRequest(BaseModel):
input: str = Field(min_length=1)
speaker: str = "Ira"
temperature: float = 0.8
top_k: int = 30
max_new_tokens: int = 2048
class TinyAya:
def __init__(self, repo_id: str, device: str):
self.device = device
self.lock = asyncio.Lock()
source = Path(repo_id)
root = source if source.exists() else Path(snapshot_download(repo_id))
dtype = torch.bfloat16 if device.startswith("cuda") else torch.float32
self.tokenizer = AutoTokenizer.from_pretrained(root, trust_remote_code=True)
self.model = AutoModelForCausalLM.from_pretrained(
root,
trust_remote_code=True,
dtype=dtype,
attn_implementation="sdpa",
).eval().to(device)
self.mimi = MimiModel.from_pretrained(root / "codec", dtype=dtype).eval().to(device)
self.sample_rate = int(AutoFeatureExtractor.from_pretrained(root / "codec").sampling_rate)
vocab = self.tokenizer.get_vocab()
self.audio_end_id = int(vocab["</audio>"])
self.mapping = {
int(token_id): (int(match.group(1)), int(match.group(2)))
for token, token_id in vocab.items()
if (match := AUDIO_RE.match(token))
}
self.allowed_ids = torch.tensor(
sorted([*self.mapping, self.audio_end_id]), device=device
)
@torch.inference_mode()
def synthesize(self, req: SpeechRequest) -> bytes:
if req.speaker not in SPEAKERS:
raise ValueError(f"speaker must be one of: {', '.join(SPEAKERS)}")
prompt = f'<text>{req.speaker}: {req.input}<audio>'
inputs = self.tokenizer(prompt, return_tensors="pt").to(self.device)
output = self.model.generate_audio(
**inputs,
allowed_ids=self.allowed_ids,
max_new_tokens=req.max_new_tokens,
min_new_tokens=8,
temperature=req.temperature,
top_k=req.top_k,
do_sample=True,
)[0].tolist()
start = inputs.input_ids.shape[1]
try:
end = output.index(self.audio_end_id, start)
except ValueError:
end = len(output)
values: list[int] = []
frame: list[int] = []
expected = 0
for token_id in output[start:end]:
item = self.mapping.get(int(token_id))
if item is None:
frame, expected = [], 0
continue
code, quantizer = item
if quantizer == expected:
frame.append(code)
expected += 1
if expected == 8:
values.extend(frame)
frame, expected = [], 0
elif quantizer == 0:
frame, expected = [code], 1
else:
frame, expected = [], 0
if not values:
raise RuntimeError("No complete Mimi-Q8 frames generated")
codes = torch.tensor(values, device=self.device).reshape(1, -1, 8).transpose(1, 2)
audio = self.mimi.decode(codes).audio_values[0, 0].float().cpu().clamp(-1, 1)
pcm = (audio.numpy() * 32767).astype("<i2")
buffer = io.BytesIO()
with wave.open(buffer, "wb") as wav:
wav.setnchannels(1)
wav.setsampwidth(2)
wav.setframerate(self.sample_rate)
wav.writeframes(pcm.tobytes())
return buffer.getvalue()
HTML = """<!doctype html><html><head><meta charset=utf-8><title>TinyAya 583</title>
<style>body{font:16px system-ui;background:#0b1020;color:#eef2ff;max-width:850px;margin:40px auto;padding:20px}textarea,select,input,button{box-sizing:border-box;width:100%;padding:12px;margin:6px 0;background:#151d33;color:#fff;border:1px solid #34405f;border-radius:6px}button{background:#f97316;border:0;font-weight:700;cursor:pointer}button:disabled{opacity:.55;cursor:wait}.row{display:grid;grid-template-columns:1fr 1fr 1fr;gap:12px}audio{width:100%;margin-top:18px}#status{min-height:24px;color:#fbbf24}</style></head>
<body><h1>TinyAya Mimi TTS 583</h1><label>Speaker</label><select id=speaker><option>Ira</option><option>Aisha</option><option>Siya</option><option>Zoya</option><option>Silver</option></select><label>Text</label><textarea id=input rows=5>&lt;description="happy, Telugu accent, fast pace"&gt; నమస్తే, ఈ రోజు ఎలా ఉన్నారు?</textarea><div class=row><input id=temp type=number value=.8 step=.1><input id=topk type=number value=30><input id=max type=number value=2048></div><button id=go onclick=run()>Synthesize</button><div id=status></div><audio id=audio controls></audio>
<script>async function run(){let b=document.getElementById('go'),s=document.getElementById('status');b.disabled=true;b.textContent='Generating...';s.textContent='Generating audio...';try{let r=await fetch('/v1/audio/speech',{method:'POST',headers:{'content-type':'application/json'},body:JSON.stringify({input:input.value,speaker:speaker.value,temperature:+temp.value,top_k:+topk.value,max_new_tokens:+max.value})});if(!r.ok)throw new Error(await r.text());audio.src=URL.createObjectURL(await r.blob());audio.play();s.textContent='Done';}catch(e){s.textContent=e.message}finally{b.disabled=false;b.textContent='Synthesize'}}</script></body></html>"""
def create_app(engine: TinyAya) -> FastAPI:
app = FastAPI(title="TinyAya 583")
@app.get("/", response_class=HTMLResponse)
def index():
return HTML
@app.get("/health")
def health():
return {"status": "ok", "model": "Pranavz/583"}
@app.post("/v1/audio/speech")
async def speech(req: SpeechRequest):
try:
async with engine.lock:
wav = await asyncio.to_thread(engine.synthesize, req)
return Response(wav, media_type="audio/wav")
except (RuntimeError, ValueError) as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return app
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--repo-id", default="Pranavz/583")
parser.add_argument("--host", default="0.0.0.0")
parser.add_argument("--port", type=int, default=7860)
parser.add_argument("--device", default="cuda")
args = parser.parse_args()
engine = TinyAya(args.repo_id, args.device)
uvicorn.run(create_app(engine), host=args.host, port=args.port)
if __name__ == "__main__":
main()