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)