# ───────────────────────────────────────────── # Mira-1-Large • ZeroGPU • OpenAI-compat API # Uses gradio.Server (extends FastAPI) so we get # ZeroGPU support + full custom POST routes. # ───────────────────────────────────────────── import os import json import time import uuid from threading import Thread import torch import spaces import gradio as gr from transformers import ( AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer, ) from fastapi import Request from fastapi.responses import StreamingResponse, JSONResponse # ── Config ──────────────────────────────────── HF_TOKEN = os.environ.get("HF_TOKEN") # set as Space secret MODEL_ID = "Bc-AI/Mira" # ← change to your private repo MODEL_KEY = "mira-1-large" # ── Load model at startup (CPU; ZeroGPU moves to GPU per-request) ───────────── print(f"[startup] loading tokenizer for {MODEL_ID} …") tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, token=HF_TOKEN) print(f"[startup] loading model for {MODEL_ID} …") model = AutoModelForCausalLM.from_pretrained( MODEL_ID, token=HF_TOKEN, torch_dtype=torch.bfloat16, device_map="cpu", # stays on CPU until @spaces.GPU kicks in low_cpu_mem_usage=True, ) model.eval() print("[startup] model ready.") # ── ZeroGPU generation kernel ───────────────── @spaces.GPU(duration=120) # raise duration for longer outputs def _generate_on_gpu(input_ids: torch.Tensor, streamer: TextIteratorStreamer, generation_kwargs: dict): """Runs entirely on GPU; ZeroGPU allocates/releases automatically.""" model.to("cuda") input_ids = input_ids.to("cuda") with torch.no_grad(): model.generate(input_ids=input_ids, streamer=streamer, **generation_kwargs) # ── Helpers ─────────────────────────────────── def _build_input_ids(messages: list) -> torch.Tensor: return tokenizer.apply_chat_template( messages, return_tensors="pt", add_generation_prompt=True, ) def _gen_kwargs(data: dict) -> dict: return { "max_new_tokens": data.get("max_tokens", 512), "temperature": data.get("temperature", 0.7), "do_sample": data.get("temperature", 0.7) > 0, "top_p": data.get("top_p", 0.95), "repetition_penalty":data.get("frequency_penalty", 1.0) + 1.0, "pad_token_id": tokenizer.eos_token_id, } def _make_chunk(content: str, model_name: str, finish: str | None = None) -> str: return "data: " + json.dumps({ "id": f"chatcmpl-{uuid.uuid4().hex}", "object": "chat.completion.chunk", "created": int(time.time()), "model": model_name, "choices": [{ "index": 0, "delta": {"content": content}, "finish_reason": finish, }], }) + "\n\n" # ── gradio.Server (FastAPI superset with ZeroGPU awareness) ─────────────────── app = gr.Server() @app.post("/v1/chat/completions") async def chat_completions(request: Request): data = await request.json() messages = data.get("messages", []) stream = data.get("stream", False) model_name = data.get("model", MODEL_KEY) kwargs = _gen_kwargs(data) input_ids = _build_input_ids(messages) # ── Streaming ───────────────────────────── if stream: streamer = TextIteratorStreamer( tokenizer, skip_prompt=True, skip_special_tokens=True ) # kick off GPU generation in a background thread thread = Thread( target=_generate_on_gpu, args=(input_ids, streamer, kwargs), daemon=True, ) thread.start() def sse_generator(): for token_text in streamer: yield _make_chunk(token_text, model_name) yield _make_chunk("", model_name, finish="stop") yield "data: [DONE]\n\n" return StreamingResponse(sse_generator(), media_type="text/event-stream") # ── Non-streaming ───────────────────────── streamer = TextIteratorStreamer( tokenizer, skip_prompt=True, skip_special_tokens=True ) thread = Thread( target=_generate_on_gpu, args=(input_ids, streamer, kwargs), daemon=True, ) thread.start() thread.join() full_text = "".join(streamer) # already drained after join prompt_tokens = input_ids.shape[-1] completion_tokens = len(tokenizer.encode(full_text, add_special_tokens=False)) return JSONResponse({ "id": f"chatcmpl-{uuid.uuid4().hex}", "object": "chat.completion", "created": int(time.time()), "model": model_name, "choices": [{ "index": 0, "message": {"role": "assistant", "content": full_text}, "finish_reason": "stop", }], "usage": { "prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens, "total_tokens": prompt_tokens + completion_tokens, }, }) @app.get("/health") async def health(): return {"status": "ok", "model": MODEL_KEY} # ── Launch ──────────────────────────────────── if __name__ == "__main__": app.launch()