acsaco's picture
Update app.py
33174d8 verified
Raw
History Blame
1.73 kB
import spaces
import gradio as gr
import torch
from fastapi import FastAPI, Request
from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer
from threading import Thread
MODEL_ID = "Qwen/Qwen3.8-27B" # Modelos >14B suelen exceder la memoria dinámica de ZeroGPU
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
torch_dtype=torch.bfloat16,
device_map="auto"
)
app = FastAPI()
@spaces.GPU(duration=120)
def generate_response(prompt: str, max_tokens: int = 2048):
inputs = tokenizer(prompt, return_tensors="pt").to("cuda")
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)
generation_kwargs = dict(inputs, streamer=streamer, max_new_tokens=max_tokens)
thread = Thread(target=model.generate, kwargs=generation_kwargs)
thread.start()
output_text = ""
for new_text in streamer:
output_text += new_text
return output_text
@app.post("/v1/chat/completions")
async def chat_completions(request: Request):
data = await request.json()
messages = data.get("messages", [])
prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
response_text = generate_response(prompt)
return {
"id": "chatcmpl-zerogpu",
"object": "chat.completion",
"choices": [{
"index": 0,
"message": {
"role": "assistant",
"content": response_text
},
"finish_reason": "stop"
}]
}
with gr.Blocks() as demo:
gr.Markdown("# Qwen ZeroGPU Endpoint")
app = gr.mount_gradio_app(app, demo, path="/")