| |
| """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() |
|
|