"""OpenAI-compatible HTTP API built on FastAPI.""" import uuid import time import json import base64 import io as _bio import threading import torch from fastapi import FastAPI, Request from fastapi.responses import JSONResponse, StreamingResponse from fastapi.middleware.cors import CORSMiddleware from models import _cached, _load_text_model, _load_image_model, _load_tts_model AVAILABLE_MODELS = { "text": [ "Qwen/Qwen2.5-3B-Instruct", "Qwen/Qwen2.5-1.5B-Instruct", "deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B", "google/flan-t5-base", "facebook/bart-large-cnn", ], "image": [ "runwayml/stable-diffusion-v1-5", ], "tts": [ "suno/bark-small", "microsoft/speecht5_tts", ], } def _openai_chat_response(model, messages, max_tokens=512, temperature=0.7, stream=True): """OpenAI-compatible chat completions (streaming or non-streaming).""" chat_id = f"chatcmpl-{uuid.uuid4().hex[:12]}" created = int(time.time()) model_obj, tokenizer, device = _cached(_load_text_model)(model) if tokenizer.chat_template: prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) else: prompt = messages[-1]["content"] if messages else "" inputs = tokenizer(prompt, return_tensors="pt").to(device) if stream: def generate(): from transformers import TextIteratorStreamer streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True) gen_kwargs = dict(max_new_tokens=max_tokens, temperature=temperature, do_sample=temperature > 0, streamer=streamer) thread = threading.Thread(target=model_obj.generate, args=(inputs.input_ids,), kwargs=gen_kwargs) thread.start() first = True for text in streamer: if text: chunk = { "id": chat_id, "object": "chat.completion.chunk", "created": created, "model": model, "choices": [{ "index": 0, "delta": {"role": "assistant", "content": text} if first else {"content": text}, "finish_reason": None, }], } first = False yield f"data: {json.dumps(chunk)}\n\n" done_chunk = { "id": chat_id, "object": "chat.completion.chunk", "created": created, "model": model, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], } yield f"data: {json.dumps(done_chunk)}\n\n" yield "data: [DONE]\n\n" return generate() else: with torch.no_grad(): output = model_obj.generate(**inputs, max_new_tokens=max_tokens, temperature=temperature, do_sample=temperature > 0) content = tokenizer.decode(output[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True) return { "id": chat_id, "object": "chat.completion", "created": created, "model": model, "choices": [{ "index": 0, "message": {"role": "assistant", "content": content}, "finish_reason": "stop", }], "usage": { "prompt_tokens": int(inputs["input_ids"].shape[1]), "completion_tokens": int(output.shape[1] - inputs["input_ids"].shape[1]), "total_tokens": int(output.shape[1]), }, } def _openai_completion_response(model, prompt, max_tokens=512, temperature=0.7, stream=True): """OpenAI-compatible text completions (streaming or non-streaming).""" comp_id = f"cmpl-{uuid.uuid4().hex[:12]}" created = int(time.time()) model_obj, tokenizer, device = _cached(_load_text_model)(model) inputs = tokenizer(prompt, return_tensors="pt").to(device) if stream: def generate(): from transformers import TextIteratorStreamer streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True) gen_kwargs = dict(max_new_tokens=max_tokens, temperature=temperature, do_sample=temperature > 0, streamer=streamer) thread = threading.Thread(target=model_obj.generate, args=(inputs.input_ids,), kwargs=gen_kwargs) thread.start() for text in streamer: if text: chunk = { "id": comp_id, "object": "text_completion", "created": created, "model": model, "choices": [{ "text": text, "index": 0, "finish_reason": None, }], } yield f"data: {json.dumps(chunk)}\n\n" done_chunk = { "id": comp_id, "object": "text_completion", "created": created, "model": model, "choices": [{"text": "", "index": 0, "finish_reason": "stop"}], } yield f"data: {json.dumps(done_chunk)}\n\n" yield "data: [DONE]\n\n" return generate() else: with torch.no_grad(): output = model_obj.generate(**inputs, max_new_tokens=max_tokens, temperature=temperature, do_sample=temperature > 0) text = tokenizer.decode(output[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True) return { "id": comp_id, "object": "text_completion", "created": created, "model": model, "choices": [{"text": text, "index": 0, "finish_reason": "stop"}], "usage": { "prompt_tokens": int(inputs["input_ids"].shape[1]), "completion_tokens": int(output.shape[1] - inputs["input_ids"].shape[1]), "total_tokens": int(output.shape[1]), }, } def _openai_image_response(model, prompt, size="512x512", steps=20, seed=42): """OpenAI-compatible image generation.""" created = int(time.time()) w, h = (int(x) for x in size.split("x")) pipe = _cached(_load_image_model)(model) gen = torch.Generator(device="cpu").manual_seed(seed) result = pipe(prompt=prompt, height=h, width=w, num_inference_steps=steps, generator=gen) buf = _bio.BytesIO() result.images[0].save(buf, format="PNG") b64 = base64.b64encode(buf.getvalue()).decode() return { "created": created, "data": [{"b64_json": b64, "revised_prompt": prompt}], } def _openai_tts_response(model, text): """OpenAI-compatible TTS (returns WAV audio bytes).""" model_obj, tokenizer, extra = _cached(_load_tts_model)(model) if tokenizer is not None: speaker = extra if extra is not None else torch.zeros((1, 512)) inputs = tokenizer(text, return_tensors="pt") with torch.no_grad(): speech = model_obj.generate(input_ids=inputs["input_ids"], speaker_embeddings=speaker) audio = speech[0].cpu().numpy() sr = 16000 else: result = model_obj(text) sr = result.get("sampling_rate", 22050) audio = result["audio"] if isinstance(audio, list): audio = audio[0] import scipy.io.wavfile as wav buf = _bio.BytesIO() wav.write(buf, sr, audio) buf.seek(0) return buf def build_fastapi_app(): """Build FastAPI app with OpenAI-compatible endpoints.""" api = FastAPI(title="HF Playground API", docs_url="/docs") api.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) @api.get("/v1/models") async def list_models(): all_models = [] for kind in ("text", "image", "tts"): for model_id in AVAILABLE_MODELS[kind]: all_models.append({ "id": model_id, "object": "model", "created": 1, "owned_by": "huggingface", "permission": [], "root": model_id, "parent": None, }) return {"object": "list", "data": all_models} @api.post("/v1/chat/completions") async def chat_completions(request: Request): try: body = await request.json() except Exception: return JSONResponse(status_code=400, content={"error": {"message": "Invalid JSON"}}) model = body.get("model", AVAILABLE_MODELS["text"][0]) messages = body.get("messages", []) max_tokens = body.get("max_tokens", 512) temperature = body.get("temperature", 0.7) stream = body.get("stream", False) if not messages: return JSONResponse(status_code=400, content={"error": {"message": "messages is required"}}) try: result = _openai_chat_response(model, messages, max_tokens, temperature, stream) if stream: return StreamingResponse(result, media_type="text/event-stream") return JSONResponse(content=result) except Exception as e: return JSONResponse(status_code=500, content={"error": {"message": str(e)}}) @api.post("/v1/completions") async def completions(request: Request): try: body = await request.json() except Exception: return JSONResponse(status_code=400, content={"error": {"message": "Invalid JSON"}}) model = body.get("model", AVAILABLE_MODELS["text"][0]) prompt = body.get("prompt", "") max_tokens = body.get("max_tokens", 512) temperature = body.get("temperature", 0.7) stream = body.get("stream", False) if not prompt: return JSONResponse(status_code=400, content={"error": {"message": "prompt is required"}}) try: result = _openai_completion_response(model, prompt, max_tokens, temperature, stream) if stream: return StreamingResponse(result, media_type="text/event-stream") return JSONResponse(content=result) except Exception as e: return JSONResponse(status_code=500, content={"error": {"message": str(e)}}) @api.post("/v1/images/generations") async def image_generations(request: Request): try: body = await request.json() except Exception: return JSONResponse(status_code=400, content={"error": {"message": "Invalid JSON"}}) model = body.get("model", AVAILABLE_MODELS["image"][0]) prompt = body.get("prompt", "") size = body.get("size", "512x512") steps = body.get("steps", 20) seed = body.get("seed", 42) if not prompt: return JSONResponse(status_code=400, content={"error": {"message": "prompt is required"}}) try: return JSONResponse(content=_openai_image_response(model, prompt, size, steps, seed)) except Exception as e: return JSONResponse(status_code=500, content={"error": {"message": str(e)}}) @api.post("/v1/audio/speech") async def audio_speech(request: Request): try: body = await request.json() except Exception: return JSONResponse(status_code=400, content={"error": {"message": "Invalid JSON"}}) model = body.get("model", AVAILABLE_MODELS["tts"][0]) text = body.get("input", "") if not text: return JSONResponse(status_code=400, content={"error": {"message": "input is required"}}) try: buf = _openai_tts_response(model, text) return StreamingResponse(iter([buf.read()]), media_type="audio/wav") except Exception as e: return JSONResponse(status_code=500, content={"error": {"message": str(e)}}) return api