reikernx commited on
Commit
3bae238
·
verified ·
1 Parent(s): 1dc7f23

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +7 -20
app.py CHANGED
@@ -1,22 +1,18 @@
1
  from flask import Flask, request, jsonify
2
  from transformers import pipeline
3
- from huggingface_hub import snapshot_download
4
  import torch, json, os
5
 
6
  app = Flask(__name__)
7
 
8
  MODEL_CACHE_DIR = os.environ.get("HF_HOME", "./model_cache")
9
 
10
- def ensure_model(model_name):
11
- local_path = snapshot_download(model_name, cache_dir=MODEL_CACHE_DIR)
12
- return local_path
13
-
14
- # Download models into cache dir
15
- path_120b = ensure_model("openai/gpt-oss-120b")
16
- path_20b = ensure_model("openai/gpt-oss-20b")
17
-
18
- pipe_120b = pipeline("text-generation", model=path_120b, torch_dtype="auto", device_map="auto")
19
- pipe_20b = pipeline("text-generation", model=path_20b, torch_dtype="auto", device_map="auto")
20
 
21
  MEMORY_FILE = "memory.json"
22
 
@@ -37,15 +33,6 @@ def generate_with_memory(pipe, session_id, user_message):
37
  save_memory()
38
  return response_text
39
 
40
- @app.route("/generate_120b")
41
- def generate_120b():
42
- msg = request.args.get("message", "")
43
- sid = request.args.get("session_id", "default")
44
- if not msg.strip():
45
- return jsonify({"error": "No message provided"}), 400
46
- resp = generate_with_memory(pipe_120b, sid, msg)
47
- return jsonify({"model": "gpt-oss-120b", "session_id": sid, "input": msg, "response": resp})
48
-
49
  @app.route("/generate_20b")
50
  def generate_20b():
51
  msg = request.args.get("message", "")
 
1
  from flask import Flask, request, jsonify
2
  from transformers import pipeline
 
3
  import torch, json, os
4
 
5
  app = Flask(__name__)
6
 
7
  MODEL_CACHE_DIR = os.environ.get("HF_HOME", "./model_cache")
8
 
9
+ # ✅ Load pre-downloaded GPT-OSS-20B model directly
10
+ pipe_20b = pipeline(
11
+ "text-generation",
12
+ model=os.path.join(MODEL_CACHE_DIR, "models--openai--gpt-oss-20b", "snapshots"),
13
+ torch_dtype="auto",
14
+ device_map="auto"
15
+ )
 
 
 
16
 
17
  MEMORY_FILE = "memory.json"
18
 
 
33
  save_memory()
34
  return response_text
35
 
 
 
 
 
 
 
 
 
 
36
  @app.route("/generate_20b")
37
  def generate_20b():
38
  msg = request.args.get("message", "")