File size: 7,036 Bytes
23ab89a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 | #!/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><description="happy, Telugu accent, fast pace"> నమస్తే, ఈ రోజు ఎలా ఉన్నారు?</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()
|