File size: 2,340 Bytes
c88ea04 eef786b c88ea04 eef786b 5d25a25 2831f7f 5d25a25 2831f7f 5d25a25 2831f7f eef786b 2831f7f 338842a 2831f7f 5d25a25 2831f7f 5d25a25 2831f7f 5d25a25 2831f7f 5d25a25 eef786b 5d25a25 2831f7f eef786b 5d25a25 2831f7f 5d25a25 2831f7f 5d25a25 2831f7f eef786b 338842a 5d25a25 c88ea04 5d25a25 c88ea04 5d25a25 c88ea04 5d25a25 c88ea04 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 | from flask import Flask, request, jsonify
from transformers import AutoModelForCausalLM, AutoTokenizer
import torch
# =========================
# 1️⃣ Load model & tokenizer
# =========================
model_name = "Qwen/Qwen2.5-0.5B-Instruct"
# Fast tokenizer for speed
tokenizer = AutoTokenizer.from_pretrained(model_name, use_fast=True)
# Load model with correct dtype
model = AutoModelForCausalLM.from_pretrained(
model_name,
device_map="auto", # Uses GPU if available, else CPU
dtype=torch.float32 # CPU inference works better with float32
)
# Optional PyTorch 2.x compile (speeds up CPU inference)
if torch.__version__.startswith("2"):
model = torch.compile(model)
# =========================
# 2️⃣ Hardcoded system prompt
# =========================
SYSTEM_PROMPT = "You are a friendly AI assistant that gives helpful and polite answers."
# =========================
# 3️⃣ Optimized chat function
# =========================
def chat(user_prompt: str):
messages = [
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": user_prompt}
]
# Apply Qwen chat template
text = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
)
# Encode input once
inputs = tokenizer([text], return_tensors="pt").to(model.device)
# Faster generation settings
outputs = model.generate(
**inputs,
max_new_tokens=128, # smaller = faster
do_sample=False, # deterministic = faster
num_beams=1 # no beam search
)
# Decode only the first sequence
response = tokenizer.decode(outputs[0], skip_special_tokens=True)
return response
# =========================
# 4️⃣ Flask app
# =========================
app = Flask(__name__)
@app.route("/chat", methods=["GET"])
def chat_route():
user_message = request.args.get("message")
if not user_message:
return jsonify({"error": "No message provided"}), 400
try:
response = chat(user_message)
return jsonify({"response": response})
except Exception as e:
return jsonify({"error": str(e)}), 500
# =========================
# 5️⃣ Run Flask
# =========================
if __name__ == "__main__":
app.run(host="0.0.0.0", port=7860) |