helloadhavan commited on
Commit
985d87b
·
1 Parent(s): daa61e1

init commit

Browse files
Files changed (1) hide show
  1. app.py +25 -41
app.py CHANGED
@@ -4,11 +4,8 @@ import os
4
 
5
  app = Flask(__name__, static_folder="static")
6
 
7
- # Put your *.gradio.live or *.hf.space URL (no trailing slash)
8
- GRADIO_SERVER_URL = os.getenv("GRADIO_SERVER_URL", "https://c9ceb67adbae7238f6.gradio.live").rstrip("/")
9
-
10
- # Default endpoint — some Spaces expose /predict by default
11
- GRADIO_PREDICT_URL = f"{GRADIO_SERVER_URL}/predict"
12
 
13
 
14
  @app.route("/")
@@ -23,52 +20,39 @@ def static_files(filename):
23
 
24
  @app.route("/generate", methods=["POST"])
25
  def generate():
26
- body = request.get_json(silent=True) or {}
27
-
28
  prompt = str(body.get("prompt", "")).strip()
29
  if not prompt:
30
  return jsonify({"error": "prompt is required"}), 400
31
 
32
- # You can optionally override if the Space uses a custom API name.
33
- # Example: body["api_name"] = "/generate" or "/my_api"
34
- api_name = body.get("api_name", "/predict")
35
- if not api_name.startswith("/"):
36
- api_name = "/" + api_name
37
-
38
- api_url = f"{GRADIO_SERVER_URL}{api_name}"
39
-
40
- # Gradio expects positional arguments if the endpoint is /predict
41
- # Usually it's a list matching the input components
42
- gradio_payload = {
43
- "data": [
44
- prompt,
45
- int(body.get("max_new_tokens", 200)),
46
- float(body.get("temperature", 0.9)),
47
- float(body.get("top_p", 0.95)),
48
- float(body.get("repetition_penalty", 1.2)),
49
- ]
50
- }
51
 
52
  try:
53
- resp = requests.post(api_url, json=gradio_payload, timeout=120)
 
 
 
 
 
 
 
54
  resp.raise_for_status()
55
- except requests.exceptions.ConnectionError:
56
- return jsonify({"error": "Cannot reach inference server."}), 502
57
- except requests.exceptions.Timeout:
58
- return jsonify({"error": "Inference server timed out."}), 504
59
- except requests.exceptions.HTTPError as e:
60
- return jsonify({"error": f"Inference server error: {str(e)}", "response": resp.text}), 502
61
 
62
- try:
63
- result = resp.json()
64
- # for many Spaces, the first item in "data" is the output
65
- generated_text = result.get("data", [None])[0] if isinstance(result, dict) else None
66
- except Exception:
67
- return jsonify({"error": "Unexpected inference server response", "raw": resp.text}), 502
68
 
69
  return jsonify({"generated_text": generated_text, "prompt": prompt})
70
 
71
 
72
  if __name__ == "__main__":
73
- port = int(os.getenv("PORT", "7860"))
74
- app.run(port=port, debug=False)
 
4
 
5
  app = Flask(__name__, static_folder="static")
6
 
7
+ # Your Gradio share link (no trailing slash)
8
+ GRADIO_SERVER_URL = os.getenv("GRADIO_SERVER_URL", "https://e35e92eb80591306dc.gradio.live").rstrip("/")
 
 
 
9
 
10
 
11
  @app.route("/")
 
20
 
21
  @app.route("/generate", methods=["POST"])
22
  def generate():
23
+ body = request.get_json(force=True) or {}
 
24
  prompt = str(body.get("prompt", "")).strip()
25
  if not prompt:
26
  return jsonify({"error": "prompt is required"}), 400
27
 
28
+ # Optional generation params with sensible defaults
29
+ max_new_tokens = int(body.get("max_new_tokens", 200))
30
+ temperature = float(body.get("temperature", 0.9))
31
+ top_p = float(body.get("top_p", 0.95))
32
+ rep_penalty = float(body.get("repetition_penalty", 1.2))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
33
 
34
  try:
35
+ resp = requests.post(
36
+ f"{GRADIO_SERVER_URL}/gradio_api/run/predict",
37
+ json={
38
+ "data": [prompt, max_new_tokens, temperature, top_p, rep_penalty],
39
+ "fn_index": 0,
40
+ },
41
+ timeout=60,
42
+ )
43
  resp.raise_for_status()
44
+ generated_text = resp.json()["data"][0]
 
 
 
 
 
45
 
46
+ except requests.exceptions.Timeout:
47
+ return jsonify({"error": "Inference timed out"}), 504
48
+ except requests.exceptions.RequestException as e:
49
+ return jsonify({"error": "Inference failed", "details": str(e)}), 502
50
+ except (KeyError, IndexError) as e:
51
+ return jsonify({"error": "Unexpected response from model server", "details": str(e)}), 502
52
 
53
  return jsonify({"generated_text": generated_text, "prompt": prompt})
54
 
55
 
56
  if __name__ == "__main__":
57
+ port = int(os.getenv("PORT", 7890))
58
+ app.run(host="0.0.0.0", port=port)