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)