spidey121 commited on
Commit
90f2cce
·
1 Parent(s): 750cd61

fix round 1

Browse files
Files changed (1) hide show
  1. 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
- # Reproducibility
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 (30%)
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 = score >= 0.5
 
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)