mdm / app.py
reikernx's picture
Update app.py
3bae238 verified
Raw
History Blame Contribute Delete
1.73 kB
from flask import Flask, request, jsonify
from transformers import pipeline
import torch, json, os
app = Flask(__name__)
MODEL_CACHE_DIR = os.environ.get("HF_HOME", "./model_cache")
# ✅ Load pre-downloaded GPT-OSS-20B model directly
pipe_20b = pipeline(
"text-generation",
model=os.path.join(MODEL_CACHE_DIR, "models--openai--gpt-oss-20b", "snapshots"),
torch_dtype="auto",
device_map="auto"
)
MEMORY_FILE = "memory.json"
def load_memory():
return json.load(open(MEMORY_FILE)) if os.path.exists(MEMORY_FILE) else {}
def save_memory():
json.dump(conversations, open(MEMORY_FILE, "w"))
conversations = load_memory()
def generate_with_memory(pipe, session_id, user_message):
conversations.setdefault(session_id, [])
conversations[session_id].append({"role": "user", "content": user_message})
outputs = pipe(user_message, max_new_tokens=256)
response_text = outputs[0]["generated_text"]
conversations[session_id].append({"role": "assistant", "content": response_text})
save_memory()
return response_text
@app.route("/generate_20b")
def generate_20b():
msg = request.args.get("message", "")
sid = request.args.get("session_id", "default")
if not msg.strip():
return jsonify({"error": "No message provided"}), 400
resp = generate_with_memory(pipe_20b, sid, msg)
return jsonify({"model": "gpt-oss-20b", "session_id": sid, "input": msg, "response": resp})
@app.route("/reset_session")
def reset_session():
sid = request.args.get("session_id", "default")
conversations.pop(sid, None)
save_memory()
return jsonify({"status": "Session reset", "session_id": sid})
if __name__ == "__main__":
app.run(host="0.0.0.0", port=7860)