| import os |
| import json |
| import asyncio |
| from fastapi import FastAPI |
| from fastapi.responses import StreamingResponse |
| from pydantic import BaseModel |
| from typing import List |
| import torch |
| from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer |
| from threading import Thread |
|
|
| app = FastAPI(title="Qwen 0.5B Streaming API") |
|
|
| model_id = "Qwen/Qwen2.5-0.5B-Instruct" |
|
|
| |
| tokenizer = AutoTokenizer.from_pretrained(model_id) |
| model = AutoModelForCausalLM.from_pretrained( |
| model_id, |
| torch_dtype=torch.float32, |
| device_map="cpu" |
| ) |
|
|
| class Message(BaseModel): |
| role: str |
| content: str |
|
|
| class ChatPayload(BaseModel): |
| messages: List[Message] |
| temperature: float = 0.7 |
| max_tokens: int = 512 |
|
|
| @app.post("/v1/chat/completions") |
| async def chat_completion(payload: ChatPayload): |
| |
| formatted_messages = [{"role": msg.role, "content": msg.content} for msg in payload.messages] |
| |
| formatted_prompt = tokenizer.apply_chat_template( |
| formatted_messages, |
| tokenize=False, |
| add_generation_prompt=True |
| ) |
| |
| inputs = tokenizer([formatted_prompt], return_tensors="pt").to("cpu") |
| |
| |
| streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True) |
| |
| generation_kwargs = dict( |
| inputs, |
| streamer=streamer, |
| max_new_tokens=payload.max_tokens, |
| do_sample=True, |
| temperature=payload.temperature, |
| pad_token_id=tokenizer.eos_token_id |
| ) |
| |
| |
| thread = Thread(target=model.generate, kwargs=generation_kwargs) |
| thread.start() |
|
|
| async def event_generator(): |
| for new_text in streamer: |
| if new_text: |
| |
| chunk = { |
| "choices": [ |
| { |
| "delta": { |
| "content": new_text |
| } |
| } |
| ] |
| } |
| yield f"data: {json.dumps(chunk)}\n\n" |
| await asyncio.sleep(0.01) |
| yield "data: [DONE]\n\n" |
|
|
| return StreamingResponse(event_generator(), media_type="text/event-stream") |
|
|
| @app.get("/") |
| async def health(): |
| return {"status": "online"} |
|
|
| if __name__ == "__main__": |
| import uvicorn |
| uvicorn.run(app, host="0.0.0.0", port=7860) |