gemma-2b-api / app.py
Singh
Update app.py
0d4c17a verified
Raw
History Blame Contribute Delete
6.66 kB
import os, json, time, threading, logging, traceback
import torch
import gradio as gr
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field
from transformers import AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer
from typing import Optional, Iterator
logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s")
logger = logging.getLogger(__name__)
MODEL_ID = os.getenv("MODEL_ID", "google/gemma-2b-it")
HF_TOKEN = os.getenv("HF_TOKEN", None)
DEVICE = "cpu"
DTYPE = torch.float32
logger.info(f"Loading {MODEL_ID} on {DEVICE} ...")
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, token=HF_TOKEN)
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
torch_dtype=DTYPE,
device_map="cpu",
token=HF_TOKEN,
)
model.eval()
logger.info("Model ready.")
class GenerateRequest(BaseModel):
prompt: str = Field(..., min_length=1, max_length=4096)
max_new_tokens: int = Field(default=128, ge=1, le=512)
temperature: float = Field(default=0.7, ge=0.01, le=2.0)
top_p: float = Field(default=0.9, ge=0.0, le=1.0)
top_k: int = Field(default=50, ge=0, le=200)
do_sample: bool = Field(default=True)
system_prompt: Optional[str] = Field(default=None, max_length=1024)
def stream_tokens(req: GenerateRequest) -> Iterator[str]:
try:
if req.system_prompt:
prompt = f"{req.system_prompt}\n\nHuman: {req.prompt}\nAssistant:"
else:
prompt = f"Human: {req.prompt}\nAssistant:"
inputs = tokenizer(
prompt, return_tensors="pt", truncation=True, max_length=1024
).to(DEVICE)
streamer = TextIteratorStreamer(
tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=120.0
)
gen_kwargs = dict(
input_ids = inputs["input_ids"],
attention_mask = inputs["attention_mask"],
streamer = streamer,
max_new_tokens = req.max_new_tokens,
temperature = req.temperature,
top_p = req.top_p,
top_k = req.top_k,
do_sample = req.do_sample,
pad_token_id = tokenizer.eos_token_id,
repetition_penalty = 1.1,
)
t = threading.Thread(target=model.generate, kwargs=gen_kwargs, daemon=True)
t.start()
token_count = 0
start = time.perf_counter()
for text in streamer:
if text:
token_count += 1
yield f"data: {json.dumps({'token': text, 'token_index': token_count})}\n\n"
t.join()
latency = (time.perf_counter() - start) * 1000
logger.info(f"Done: {token_count} tokens in {latency:.0f}ms")
yield f"data: {json.dumps({'done': True, 'total_tokens': token_count, 'latency_ms': round(latency, 1)})}\n\n"
except Exception as e:
tb = traceback.format_exc()
logger.error(tb)
yield f"data: {json.dumps({'error': str(e), 'traceback': tb})}\n\n"
# ── FastAPI app ──
api = FastAPI(title="OPT-350M API")
api.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=False,
allow_methods=["*"],
allow_headers=["*"],
expose_headers=["*"],
)
@api.get("/api/health")
async def health():
return {
"status" : "ok",
"model" : MODEL_ID,
"device" : DEVICE,
"dtype" : str(DTYPE),
}
@api.post("/api/generate/stream")
async def generate_stream(req: GenerateRequest):
return StreamingResponse(
stream_tokens(req),
media_type="text/event-stream",
headers={
"Cache-Control" : "no-cache",
"X-Accel-Buffering" : "no",
"Access-Control-Allow-Origin": "*",
},
)
@api.post("/api/generate")
async def generate(req: GenerateRequest):
full_text = ""
total_tokens = 0
latency_ms = 0.0
for chunk in stream_tokens(req):
if not chunk.startswith("data: "):
continue
try:
data = json.loads(chunk[6:])
except json.JSONDecodeError:
continue
if "error" in data:
raise HTTPException(status_code=500, detail=data["error"])
if "token" in data:
full_text += data["token"]
total_tokens += 1
if "done" in data:
latency_ms = data["latency_ms"]
return {
"generated_text" : full_text,
"completion_tokens" : total_tokens,
"latency_ms" : latency_ms,
"model" : MODEL_ID,
}
# ── Gradio UI ──
def gradio_generate(prompt, system_prompt, max_new_tokens, temperature):
req = GenerateRequest(
prompt = prompt,
system_prompt = system_prompt or None,
max_new_tokens = int(max_new_tokens),
temperature = temperature,
)
result = ""
for chunk in stream_tokens(req):
if chunk.startswith("data: "):
try:
data = json.loads(chunk[6:])
if "token" in data:
result += data["token"]
yield result
except json.JSONDecodeError:
pass
with gr.Blocks(title="OPT-350M API") as demo:
gr.Markdown(
"## OPT-350M β€” Streaming API\n"
"Use `/api/generate/stream` or `/api/generate` from your backend.\n\n"
"**Health check:** `/api/health`"
)
with gr.Row():
with gr.Column():
sys_box = gr.Textbox(label="System prompt (optional)", lines=2)
prompt_box = gr.Textbox(label="Prompt", lines=4, placeholder="Ask something...")
with gr.Row():
max_tok = gr.Slider(32, 512, value=128, step=32, label="Max tokens")
temp = gr.Slider(0.01, 2.0, value=0.7, step=0.05, label="Temperature")
btn = gr.Button("Generate", variant="primary")
with gr.Column():
output = gr.Textbox(label="Output", lines=12)
btn.click(fn=gradio_generate, inputs=[prompt_box, sys_box, max_tok, temp], outputs=output)
# ── Mount FastAPI routes into Gradio and launch ──
# This is the correct way to keep the app alive on HF Spaces
app = gr.mount_gradio_app(api, demo, path="/")
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=7860)