Mira-API-free / app.py
Bc-AI's picture
Create app.py
3bb28fd verified
Raw
History Blame Contribute Delete
5.86 kB
import gradio as gr
from gradio import Server
import spaces
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM
# ─────────────────────────────────────────────
# 1. MODEL SETUP
# ─────────────────────────────────────────────
MODEL_ID = "Smilyai-labs/Mira-1-large"
print(f"Loading tokenizer: {MODEL_ID}")
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, trust_remote_code=True)
print(f"Loading model: {MODEL_ID}")
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
torch_dtype=torch.float16,
device_map="auto", # ZeroGPU manages CUDA device
trust_remote_code=True, # Needed for Qwen-based custom archs
)
model.eval()
# ─────────────────────────────────────────────
# 2. INFERENCE FUNCTION (ZeroGPU decorated)
# ─────────────────────────────────────────────
@spaces.GPU(duration=120)
def generate(
prompt: str,
system_prompt: str = "You are a helpful assistant.",
max_new_tokens: int = 512,
temperature: float = 0.7,
top_p: float = 0.9,
do_sample: bool = True,
) -> str:
"""
Generate a text response from Mira-1-Large.
Args:
prompt: The user message / prompt to send to the model.
system_prompt: System-level instruction for the model.
max_new_tokens: Maximum number of tokens to generate.
temperature: Sampling temperature (higher = more creative).
top_p: Nucleus sampling probability mass.
do_sample: Whether to use sampling (True) or greedy decoding (False).
Returns:
The model's text response as a string.
"""
# Build chat-style messages (Qwen uses apply_chat_template)
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": prompt},
]
# Qwen / Mira chat template
text = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
inputs = tokenizer(text, return_tensors="pt").to(model.device)
with torch.no_grad():
output_ids = model.generate(
**inputs,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
do_sample=do_sample,
pad_token_id=tokenizer.eos_token_id,
)
# Decode only the newly generated tokens
new_tokens = output_ids[0][inputs["input_ids"].shape[1]:]
response = tokenizer.decode(new_tokens, skip_special_tokens=True)
return response
# ─────────────────────────────────────────────
# 3. STREAMING INFERENCE (SSE / token-by-token)
# ─────────────────────────────────────────────
@spaces.GPU(duration=120)
def generate_stream(
prompt: str,
system_prompt: str = "You are a helpful assistant.",
max_new_tokens: int = 512,
temperature: float = 0.7,
top_p: float = 0.9,
) -> str:
"""
Stream a text response token-by-token from Mira-1-Large via SSE.
Args:
prompt: The user message / prompt.
system_prompt: System-level instruction for the model.
max_new_tokens: Maximum number of tokens to generate.
temperature: Sampling temperature.
top_p: Nucleus sampling probability mass.
Yields:
Partial response strings, growing with each new token.
"""
from transformers import TextIteratorStreamer
from threading import Thread
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": prompt},
]
text = tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
inputs = tokenizer(text, return_tensors="pt").to(model.device)
streamer = TextIteratorStreamer(
tokenizer, skip_prompt=True, skip_special_tokens=True
)
gen_kwargs = dict(
**inputs,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
do_sample=True,
streamer=streamer,
pad_token_id=tokenizer.eos_token_id,
)
thread = Thread(target=model.generate, kwargs=gen_kwargs)
thread.start()
partial = ""
for new_text in streamer:
partial += new_text
yield partial
# ─────────────────────────────────────────────
# 4. gr.Server β€” REST API + OPTIONAL SWAGGER UI
# ─────────────────────────────────────────────
app = Server(
title="Mira-1-Large API",
summary="ZeroGPU-backed REST API for Smilyai-labs/Mira-1-large (Qwen arch)",
version="1.0.0",
)
# Register as Gradio API endpoints (queued, SSE-streaming capable)
app.api(generate, name="generate") # POST /gradio_api/call/generate
app.api(generate_stream, name="generate_stream") # POST /gradio_api/call/generate_stream
# Optional: plain FastAPI GET health-check route
@app.get("/health")
def health():
return {"status": "ok", "model": MODEL_ID}
# ─────────────────────────────────────────────
# 5. LAUNCH
# ─────────────────────────────────────────────
app.launch()