Lachowski-V2's picture
Update app.py
7f717aa verified
Raw
History Blame Contribute Delete
2.61 kB
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"
# Initialize model directly onto CPU
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):
# Convert incoming array to ChatML format
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")
# Initialize the background streamer tool
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
)
# Run text generation in a separate thread to keep the main event loop unblocked
thread = Thread(target=model.generate, kwargs=generation_kwargs)
thread.start()
async def event_generator():
for new_text in streamer:
if new_text:
# Mirror standard OpenAI stream object format
chunk = {
"choices": [
{
"delta": {
"content": new_text
}
}
]
}
yield f"data: {json.dumps(chunk)}\n\n"
await asyncio.sleep(0.01) # Small yield breather for network loop
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)