| 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" |
|
|
|
|
| |
| 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, |
| } |
|
|
|
|
| |
| 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) |
|
|
|
|
| |
| |
| app = gr.mount_gradio_app(api, demo, path="/") |
|
|
| if __name__ == "__main__": |
| import uvicorn |
| uvicorn.run(app, host="0.0.0.0", port=7860) |