Spaces:
Sleeping
Sleeping
| 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 1.5B Streaming API") | |
| model_id = "Qwen/Qwen2.5-1.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 | |
| 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") | |
| async def health(): | |
| return {"status": "online"} | |
| if __name__ == "__main__": | |
| import uvicorn | |
| uvicorn.run(app, host="0.0.0.0", port=7860) |