Spaces:
Sleeping
Sleeping
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # 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 βββββββββββββββββ | |
| # 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() | |
| 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, | |
| }, | |
| }) | |
| async def health(): | |
| return {"status": "ok", "model": MODEL_KEY} | |
| # ββ Launch ββββββββββββββββββββββββββββββββββββ | |
| if __name__ == "__main__": | |
| app.launch() |