Spaces:
Sleeping
Sleeping
Commit Β·
c8e7cd5
1
Parent(s): 6025d0d
Final commit :
Browse files- inference.py +5 -13
- openenv.yaml +0 -1
- requirements.txt +2 -3
inference.py
CHANGED
|
@@ -14,9 +14,6 @@ import sys
|
|
| 14 |
import json
|
| 15 |
import time
|
| 16 |
import requests
|
| 17 |
-
from openai import OpenAI
|
| 18 |
-
from dotenv import load_dotenv
|
| 19 |
-
load_dotenv()
|
| 20 |
|
| 21 |
# ββ Configuration βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 22 |
|
|
@@ -28,7 +25,6 @@ API_BASE_URL = os.environ.get("API_BASE_URL", "https://api.groq.com/openai/v1")
|
|
| 28 |
MODEL_NAME = os.environ.get("MODEL_NAME", "llama-3.1-8b-instant")
|
| 29 |
HF_TOKEN = os.environ.get("HF_TOKEN")
|
| 30 |
|
| 31 |
-
# client is initialized inside main() to avoid startup crashes
|
| 32 |
client = None
|
| 33 |
|
| 34 |
# ββ Stdout log functions (mandatory format) βββββββββββββββββββββββββββββββββββ
|
|
@@ -46,8 +42,6 @@ def log_end(success, steps, rewards):
|
|
| 46 |
rewards_str = ",".join(f"{r:.2f}" for r in rewards)
|
| 47 |
print(f"[END] success={str(success).lower()} steps={steps} rewards={rewards_str}", flush=True)
|
| 48 |
|
| 49 |
-
# ββ Debug log (stderr only) βββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 50 |
-
|
| 51 |
def debug(msg):
|
| 52 |
print(msg, file=sys.stderr, flush=True)
|
| 53 |
|
|
@@ -140,7 +134,7 @@ def ask_llm(task_description, schema, hint, attempt, previous_attempts):
|
|
| 140 |
# ββ Task solver βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 141 |
|
| 142 |
def solve_task(task_id):
|
| 143 |
-
task_name
|
| 144 |
|
| 145 |
try:
|
| 146 |
reset_resp = env_reset(task_id)
|
|
@@ -239,11 +233,11 @@ def solve_task(task_id):
|
|
| 239 |
def main():
|
| 240 |
global client
|
| 241 |
|
| 242 |
-
# Initialize client inside main() to avoid import-time crashes
|
| 243 |
try:
|
| 244 |
-
|
|
|
|
| 245 |
client = OpenAI(
|
| 246 |
-
base_url=
|
| 247 |
api_key=api_key,
|
| 248 |
)
|
| 249 |
except Exception as e:
|
|
@@ -255,7 +249,6 @@ def main():
|
|
| 255 |
debug(f"API Base : {API_BASE_URL}")
|
| 256 |
debug(f"Env Server : {ENV_BASE_URL}")
|
| 257 |
|
| 258 |
-
# Get task_id from command line or environment variable
|
| 259 |
if len(sys.argv) > 1:
|
| 260 |
task_id = int(sys.argv[1])
|
| 261 |
else:
|
|
@@ -265,8 +258,7 @@ def main():
|
|
| 265 |
|
| 266 |
result = solve_task(task_id)
|
| 267 |
|
| 268 |
-
debug(f"\
|
| 269 |
-
debug(f"RESULT: Task {result['task_id']} ({result['difficulty']}) β {'SOLVED' if result['solved'] else f'best={result[chr(98)+chr(101)+chr(115)+chr(116)+chr(95)+chr(114)+chr(101)+chr(119)+chr(97)+chr(114)+chr(100)]:.3f}'}")
|
| 270 |
|
| 271 |
with open("results.json", "w") as f:
|
| 272 |
json.dump({
|
|
|
|
| 14 |
import json
|
| 15 |
import time
|
| 16 |
import requests
|
|
|
|
|
|
|
|
|
|
| 17 |
|
| 18 |
# ββ Configuration βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 19 |
|
|
|
|
| 25 |
MODEL_NAME = os.environ.get("MODEL_NAME", "llama-3.1-8b-instant")
|
| 26 |
HF_TOKEN = os.environ.get("HF_TOKEN")
|
| 27 |
|
|
|
|
| 28 |
client = None
|
| 29 |
|
| 30 |
# ββ Stdout log functions (mandatory format) βββββββββββββββββββββββββββββββββββ
|
|
|
|
| 42 |
rewards_str = ",".join(f"{r:.2f}" for r in rewards)
|
| 43 |
print(f"[END] success={str(success).lower()} steps={steps} rewards={rewards_str}", flush=True)
|
| 44 |
|
|
|
|
|
|
|
| 45 |
def debug(msg):
|
| 46 |
print(msg, file=sys.stderr, flush=True)
|
| 47 |
|
|
|
|
| 134 |
# ββ Task solver βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 135 |
|
| 136 |
def solve_task(task_id):
|
| 137 |
+
task_name = f"sql-task-{task_id}"
|
| 138 |
|
| 139 |
try:
|
| 140 |
reset_resp = env_reset(task_id)
|
|
|
|
| 233 |
def main():
|
| 234 |
global client
|
| 235 |
|
|
|
|
| 236 |
try:
|
| 237 |
+
from openai import OpenAI
|
| 238 |
+
api_key = os.environ.get("HF_TOKEN") or os.environ.get("API_KEY") or "no-key-needed"
|
| 239 |
client = OpenAI(
|
| 240 |
+
base_url=API_BASE_URL,
|
| 241 |
api_key=api_key,
|
| 242 |
)
|
| 243 |
except Exception as e:
|
|
|
|
| 249 |
debug(f"API Base : {API_BASE_URL}")
|
| 250 |
debug(f"Env Server : {ENV_BASE_URL}")
|
| 251 |
|
|
|
|
| 252 |
if len(sys.argv) > 1:
|
| 253 |
task_id = int(sys.argv[1])
|
| 254 |
else:
|
|
|
|
| 258 |
|
| 259 |
result = solve_task(task_id)
|
| 260 |
|
| 261 |
+
debug(f"\nRESULT: Task {result['task_id']} ({result['difficulty']}) β {'SOLVED' if result['solved'] else 'best=' + str(round(result['best_reward'], 3))}")
|
|
|
|
| 262 |
|
| 263 |
with open("results.json", "w") as f:
|
| 264 |
json.dump({
|
openenv.yaml
CHANGED
|
@@ -160,5 +160,4 @@ inference:
|
|
| 160 |
llm_env_vars:
|
| 161 |
- API_BASE_URL
|
| 162 |
- MODEL_NAME
|
| 163 |
-
- API_KEY
|
| 164 |
- HF_TOKEN
|
|
|
|
| 160 |
llm_env_vars:
|
| 161 |
- API_BASE_URL
|
| 162 |
- MODEL_NAME
|
|
|
|
| 163 |
- HF_TOKEN
|
requirements.txt
CHANGED
|
@@ -2,7 +2,6 @@ fastapi==0.115.0
|
|
| 2 |
uvicorn==0.30.6
|
| 3 |
pydantic==2.9.2
|
| 4 |
requests==2.32.3
|
| 5 |
-
openai=
|
| 6 |
-
python-dotenv==1.0.1
|
| 7 |
pyyaml==6.0.2
|
| 8 |
-
|
|
|
|
| 2 |
uvicorn==0.30.6
|
| 3 |
pydantic==2.9.2
|
| 4 |
requests==2.32.3
|
| 5 |
+
openai>=1.0.0,<2.0.0
|
|
|
|
| 6 |
pyyaml==6.0.2
|
| 7 |
+
python-dotenv
|