inf-end-1 / app.py
Bc-AI's picture
Create app.py
bb3e531 verified
Raw
History Blame Contribute Delete
5.94 kB
# ─────────────────────────────────────────────
# 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()