Spaces:
Runtime error
Runtime error
fix round 1
Browse files- inference.py +32 -9
inference.py
CHANGED
|
@@ -1,15 +1,12 @@
|
|
| 1 |
-
import threading
|
| 2 |
import time
|
| 3 |
import os
|
| 4 |
import random
|
| 5 |
import traceback
|
|
|
|
|
|
|
| 6 |
from openai import OpenAI
|
| 7 |
|
| 8 |
-
|
| 9 |
-
random.seed(42)
|
| 10 |
-
|
| 11 |
-
# Start fake server
|
| 12 |
-
|
| 13 |
from env.env import DeceptionEnv
|
| 14 |
from env.attacker import simulate_attack
|
| 15 |
|
|
@@ -18,6 +15,31 @@ from tasks.easy.grader import grade as easy_grade
|
|
| 18 |
from tasks.medium.grader import grade as medium_grade
|
| 19 |
from tasks.hard.grader import grade as hard_grade
|
| 20 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 21 |
|
| 22 |
# Environment variables
|
| 23 |
API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")
|
|
@@ -103,7 +125,7 @@ def run_task(task_name):
|
|
| 103 |
# AI chooses action
|
| 104 |
action = choose_action(client, state)
|
| 105 |
|
| 106 |
-
# exploration
|
| 107 |
if random.random() < 0.3:
|
| 108 |
action = random.choice(env.action_space())
|
| 109 |
|
|
@@ -134,7 +156,8 @@ def run_task(task_name):
|
|
| 134 |
else:
|
| 135 |
score = hard_grade(rewards)
|
| 136 |
|
| 137 |
-
success
|
|
|
|
| 138 |
|
| 139 |
print(
|
| 140 |
f"[END] success={str(success).lower()} steps={steps} "
|
|
@@ -156,4 +179,4 @@ run_task("medium")
|
|
| 156 |
run_task("hard")
|
| 157 |
|
| 158 |
# Keep alive briefly
|
| 159 |
-
time.sleep(120)
|
|
|
|
|
|
|
| 1 |
import time
|
| 2 |
import os
|
| 3 |
import random
|
| 4 |
import traceback
|
| 5 |
+
import requests
|
| 6 |
+
import threading
|
| 7 |
from openai import OpenAI
|
| 8 |
|
| 9 |
+
from env.fake_server import run_server
|
|
|
|
|
|
|
|
|
|
|
|
|
| 10 |
from env.env import DeceptionEnv
|
| 11 |
from env.attacker import simulate_attack
|
| 12 |
|
|
|
|
| 15 |
from tasks.medium.grader import grade as medium_grade
|
| 16 |
from tasks.hard.grader import grade as hard_grade
|
| 17 |
|
| 18 |
+
# Reproducibility
|
| 19 |
+
random.seed(42)
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
# Wait for server to start
|
| 23 |
+
def wait_for_server():
|
| 24 |
+
for _ in range(15):
|
| 25 |
+
try:
|
| 26 |
+
r = requests.get("http://127.0.0.1:7860/status")
|
| 27 |
+
if r.status_code == 200:
|
| 28 |
+
print("Server ready")
|
| 29 |
+
return
|
| 30 |
+
except:
|
| 31 |
+
pass
|
| 32 |
+
time.sleep(1)
|
| 33 |
+
|
| 34 |
+
raise RuntimeError("Server not started")
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
# Start fake server
|
| 38 |
+
threading.Thread(target=run_server, daemon=True).start()
|
| 39 |
+
|
| 40 |
+
# Wait for server
|
| 41 |
+
wait_for_server()
|
| 42 |
+
|
| 43 |
|
| 44 |
# Environment variables
|
| 45 |
API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")
|
|
|
|
| 125 |
# AI chooses action
|
| 126 |
action = choose_action(client, state)
|
| 127 |
|
| 128 |
+
# exploration
|
| 129 |
if random.random() < 0.3:
|
| 130 |
action = random.choice(env.action_space())
|
| 131 |
|
|
|
|
| 156 |
else:
|
| 157 |
score = hard_grade(rewards)
|
| 158 |
|
| 159 |
+
# success threshold
|
| 160 |
+
success = score >= 0.3
|
| 161 |
|
| 162 |
print(
|
| 163 |
f"[END] success={str(success).lower()} steps={steps} "
|
|
|
|
| 179 |
run_task("hard")
|
| 180 |
|
| 181 |
# Keep alive briefly
|
| 182 |
+
time.sleep(120)
|