Skip to content

Commit 78b00cb

Browse files
committed
debug: (inference) yeah, given up fully
1 parent 85e5b32 commit 78b00cb

1 file changed

Lines changed: 125 additions & 205 deletions

File tree

inference.py

Lines changed: 125 additions & 205 deletions
Original file line numberDiff line numberDiff line change
@@ -7,265 +7,185 @@
77
import threading
88
from pathlib import Path
99

10-
try :
11-
import httpx
12-
from dotenv import load_dotenv
13-
from openai import OpenAI
14-
except Exception as e:
15-
print("WHATTTTT : ", e)
16-
sys.exit(0)
10+
import httpx
11+
from dotenv import load_dotenv
12+
from openai import OpenAI
1713

18-
try:
19-
ROOT = Path(__file__).parent.resolve()
20-
if str(ROOT) not in sys.path:
21-
sys.path.insert(0, str(ROOT))
14+
ROOT = Path(__file__).parent.resolve()
15+
if str(ROOT) not in sys.path:
16+
sys.path.insert(0, str(ROOT))
2217

23-
except Exception as e:
24-
print("BRUH : ", e)
25-
sys.exit(0)
18+
TASK_DIR = ROOT / "tasks"
2619

20+
from cloudenv.client import ClientEnvironment # noqa
21+
from cloudenv.models import CloudEnvAction, CloudEnvObservation # noqa
2722

28-
try:
29-
TASK_DIR = ROOT / "tasks"
30-
from cloudenv.client import ClientEnvironment
31-
from cloudenv.models import CloudEnvAction, CloudEnvObservation
23+
load_dotenv()
3224

33-
load_dotenv()
25+
HF_TOKEN = os.getenv("HF_TOKEN") or os.getenv("API_KEY")
26+
API_BASE_URL = os.getenv("API_BASE_URL", "https://api.openai.com/v1")
27+
MODEL_NAME = os.getenv("MODEL_NAME", "gpt-4o")
3428

35-
HF_TOKEN = os.getenv("HF_TOKEN") or os.getenv("API_KEY")
36-
API_BASE_URL = os.getenv("API_BASE_URL", "https://api.openai.com/v1")
37-
MODEL_NAME = os.getenv("MODEL_NAME", "gpt-4o")
29+
MAX_STEPS = int(os.getenv("MAX_STEPS", "20"))
30+
TEMPERATURE = float(os.getenv("TEMPERATURE", "0.3"))
31+
MAX_TOKENS = int(os.getenv("MAX_TOKENS", "300"))
3832

39-
MAX_STEPS = int(os.getenv("MAX_STEPS", "20"))
40-
TEMPERATURE = float(os.getenv("TEMPERATURE", "0.3"))
41-
MAX_TOKENS = int(os.getenv("MAX_TOKENS", "300"))
33+
SERVER_URL = "http://127.0.0.1:8000"
4234

43-
SERVER_URL = "http://127.0.0.1:8000"
35+
TASK_ORDER = ["easy-task", "medium-task", "hard-task"]
4436

45-
TASK_ORDER = ["easy-task", "medium-task", "hard-task"]
37+
client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN)
4638

47-
client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN)
48-
SYSTEM_PROMPT = 'Return ONLY JSON: {"service":"","operation":"","instance_id":null,"payload":{}}'
49-
except Exception as e:
50-
print("At this point, I give up : ", e)
51-
sys.exit(0)
39+
SYSTEM_PROMPT = '{"service":"","operation":"","instance_id":null,"payload":{}}'
5240

5341

5442
def start_ministack():
55-
try:
56-
proc = subprocess.Popen(
57-
["uv", "run", "ministack"],
58-
cwd=ROOT,
59-
stdout=subprocess.DEVNULL,
60-
stderr=subprocess.DEVNULL,
61-
)
62-
time.sleep(3)
63-
return proc
64-
except Exception as e:
65-
print(f"[MINISTACK START ERROR] {e}", flush=True)
66-
raise
43+
proc = subprocess.Popen(
44+
["uv", "run", "ministack"],
45+
cwd=ROOT,
46+
stdout=subprocess.DEVNULL,
47+
stderr=subprocess.DEVNULL,
48+
)
49+
time.sleep(3)
50+
return proc
6751

6852

6953
def start_server():
70-
try:
71-
proc = subprocess.Popen(
72-
["uv", "run", "server"],
73-
cwd=ROOT,
74-
stdout=subprocess.PIPE,
75-
stderr=subprocess.STDOUT,
76-
text=True,
77-
bufsize=1,
78-
)
54+
proc = subprocess.Popen(
55+
["uv", "run", "server"],
56+
cwd=ROOT,
57+
stdout=subprocess.PIPE,
58+
stderr=subprocess.STDOUT,
59+
text=True,
60+
bufsize=1,
61+
)
7962

80-
def stream():
81-
try:
82-
for line in proc.stdout:
83-
print(f"[SERVER] {line}", end="")
84-
except Exception as e:
85-
print(f"[SERVER STREAM ERROR] {e}", flush=True)
63+
def stream():
64+
for line in proc.stdout:
65+
print(f"[SERVER] {line}", end="")
8666

87-
threading.Thread(target=stream, daemon=True).start()
67+
threading.Thread(target=stream, daemon=True).start()
8868

89-
for _ in range(40):
90-
try:
91-
r = httpx.get(f"{SERVER_URL}/docs", timeout=1)
92-
if r.status_code == 200:
93-
print("[INFO] server ready", flush=True)
94-
return proc
95-
except Exception:
96-
time.sleep(1)
69+
for _ in range(40):
70+
r = httpx.get(f"{SERVER_URL}/docs", timeout=1)
71+
if r.status_code == 200:
72+
print("[INFO] server ready")
73+
return proc
74+
time.sleep(1)
9775

98-
raise RuntimeError("server failed to start")
99-
100-
except Exception as e:
101-
print(f"[SERVER START ERROR] {e}", flush=True)
102-
raise
76+
raise RuntimeError("server failed")
10377

10478

10579
def load_task(name: str):
106-
try:
107-
return json.loads((TASK_DIR / f"{name}.json").read_text())
108-
except Exception as e:
109-
print(f"[TASK LOAD ERROR] {name}: {e}", flush=True)
110-
raise
80+
return json.loads((TASK_DIR / f"{name}.json").read_text())
11181

11282

11383
def register_pipeline(task):
114-
try:
115-
url = f"{SERVER_URL}/api/pipelines"
116-
with httpx.Client(timeout=20) as http:
117-
r = http.post(url, json=task)
84+
url = f"{SERVER_URL}/api/pipelines"
85+
with httpx.Client(timeout=20) as http:
86+
r = http.post(url, json=task)
11887

119-
if r.status_code == 400:
120-
try:
121-
detail = r.json().get("detail", "")
122-
if "already exists" in detail:
123-
http.delete(f"{url}/{task['pipeline_id']}")
124-
http.post(url, json=task)
125-
except Exception:
126-
pass
88+
if r.status_code == 400:
89+
detail = r.json().get("detail", "")
90+
if "already exists" in detail:
91+
http.delete(f"{url}/{task['pipeline_id']}")
92+
http.post(url, json=task)
12793

128-
r.raise_for_status()
129-
except Exception as e:
130-
print(f"[REGISTER ERROR] {e}", flush=True)
94+
r.raise_for_status()
13195

13296

13397
def parse_action(raw: str, task_id: str, episode_id: str) -> CloudEnvAction:
134-
try:
135-
raw = raw.strip()
136-
raw = raw[raw.find("{") : raw.rfind("}") + 1]
137-
d = json.loads(raw)
138-
return CloudEnvAction(
139-
task_id=task_id,
140-
episode_id=episode_id,
141-
service=d.get("service", "internal"),
142-
instance_id=d.get("instance_id"),
143-
operation=d.get("operation", "GetCallerIdentity"),
144-
payload=d.get("payload", {}),
145-
principal_arn=d.get("principal_arn"),
146-
intent=d.get("intent"),
147-
)
148-
except Exception:
149-
return CloudEnvAction(
150-
task_id=task_id,
151-
episode_id=episode_id,
152-
service="internal",
153-
operation="GetCallerIdentity",
154-
payload={},
155-
)
98+
raw = raw.strip()
99+
raw = raw[raw.find("{") : raw.rfind("}") + 1]
100+
d = json.loads(raw)
101+
102+
return CloudEnvAction(
103+
task_id=task_id,
104+
episode_id=episode_id,
105+
service=d.get("service", "internal"),
106+
instance_id=d.get("instance_id"),
107+
operation=d.get("operation", "GetCallerIdentity"),
108+
payload=d.get("payload", {}),
109+
principal_arn=d.get("principal_arn"),
110+
intent=d.get("intent"),
111+
)
156112

157113

158114
def call_model(step, state, broken, ideal, history):
159-
try:
160-
prompt = f"""
115+
prompt = f"""
161116
STEP {step}
162117
STATE {json.dumps(state)[:1000]}
163118
BROKEN {json.dumps(broken)[:1000]}
164119
IDEAL {json.dumps(ideal)[:1000]}
165120
HISTORY {history[-5:]}
166121
"""
167-
res = client.chat.completions.create(
168-
model=MODEL_NAME,
169-
messages=[
170-
{"role": "system", "content": SYSTEM_PROMPT},
171-
{"role": "user", "content": prompt},
172-
],
173-
temperature=TEMPERATURE,
174-
max_tokens=MAX_TOKENS,
175-
)
176-
return (res.choices[0].message.content or "").strip()
177-
except Exception as e:
178-
print(f"[MODEL ERROR] {e}", flush=True)
179-
return "{}"
180-
181-
182-
async def run_episode(task):
183-
try:
184-
pipeline_id = task["pipeline_id"]
185-
task_id = pipeline_id
186-
episode_id = f"{pipeline_id}-episode-1"
187-
188-
env = ClientEnvironment(base_url=SERVER_URL)
189-
190-
try:
191-
register_pipeline(task)
192-
except Exception as e:
193-
print(f"[REGISTER ERROR] {e}", flush=True)
194-
195-
try:
196-
result = await env.initialize_environment(
197-
task_id=task_id,
198-
episode_id=episode_id,
199-
)
200-
except Exception as e:
201-
print(f"[INIT ERROR] {e}", flush=True)
202-
return
203122

204-
obs = result.observation
123+
res = client.chat.completions.create(
124+
model=MODEL_NAME,
125+
messages=[
126+
{"role": "system", "content": SYSTEM_PROMPT},
127+
{"role": "user", "content": prompt},
128+
],
129+
temperature=TEMPERATURE,
130+
max_tokens=MAX_TOKENS,
131+
)
205132

206-
history = []
207-
rewards = []
133+
return (res.choices[0].message.content or "").strip()
208134

209-
for step in range(MAX_STEPS):
210-
try:
211-
if obs.done:
212-
break
213135

214-
raw = call_model(
215-
step,
216-
obs.current_pipeline_state,
217-
task["broken_pipeline"],
218-
task["ideal_pipeline"],
219-
history,
220-
)
136+
def run_episode(task):
137+
pipeline_id = task["pipeline_id"]
138+
task_id = pipeline_id
139+
episode_id = f"{pipeline_id}-episode-1"
221140

222-
action = parse_action(raw, task_id, episode_id)
141+
env = ClientEnvironment(base_url=SERVER_URL)
223142

224-
result = await env.step(action)
225-
obs = result.observation
143+
register_pipeline(task)
226144

227-
rewards.append(float(obs.reward or 0))
228-
history.append(raw)
145+
result = asyncio.run(
146+
env.initialize_environment(task_id=task_id, episode_id=episode_id)
147+
)
229148

230-
except Exception as e:
231-
print(f"[STEP ERROR] {e}", flush=True)
232-
break
149+
obs = result.observation
150+
151+
history = []
152+
rewards = []
153+
154+
for step in range(MAX_STEPS):
155+
if obs.done:
156+
break
157+
158+
raw = call_model(
159+
step,
160+
obs.current_pipeline_state,
161+
task["broken_pipeline"],
162+
task["ideal_pipeline"],
163+
history,
164+
)
165+
166+
action = parse_action(raw, task_id, episode_id)
167+
168+
result = asyncio.run(env.step(action))
169+
obs = result.observation
233170

234-
score = sum(rewards) / (len(rewards) or 1)
235-
print(f"[END] task={task_id} score={score:.3f}", flush=True)
171+
rewards.append(float(obs.reward or 0))
172+
history.append(raw)
236173

237-
except Exception as e:
238-
print(f"[EPISODE ERROR] {e}", flush=True)
174+
score = sum(rewards) / (len(rewards) or 1)
175+
print(f"[END] task={task_id} score={score:.3f}")
239176

240177

241-
async def main():
242-
try:
243-
miniproc = start_ministack()
244-
serverproc = start_server()
178+
def main():
179+
miniproc = start_ministack()
180+
serverproc = start_server()
245181

246-
try:
247-
for name in TASK_ORDER:
248-
try:
249-
task = load_task(name)
250-
await run_episode(task)
251-
except Exception as e:
252-
print(f"[TASK ERROR] {name}: {e}", flush=True)
253-
finally:
254-
try:
255-
miniproc.terminate()
256-
except Exception:
257-
pass
258-
try:
259-
serverproc.terminate()
260-
except Exception:
261-
pass
182+
for name in TASK_ORDER:
183+
task = load_task(name)
184+
run_episode(task)
262185

263-
except Exception as e:
264-
print(f"[FATAL ERROR] {e}", flush=True)
186+
miniproc.terminate()
187+
serverproc.terminate()
265188

266189

267190
if __name__ == "__main__":
268-
try:
269-
asyncio.run(main())
270-
except Exception as e:
271-
print(f"[UNHANDLED TOP LEVEL ERROR] {e}", flush=True)
191+
main()

0 commit comments

Comments
 (0)