sql-query-env / inference.py
Deepikachintamreddy
SQL Query OpenEnv Environment
ede2fa4
Raw
History Blame Contribute Delete
6.07 kB
"""
Baseline inference script for the SQL Query Environment.
Uses the OpenAI-compatible API to run an LLM agent against all 3 tasks.
Reads credentials from environment variables:
- API_BASE_URL: The API endpoint for the LLM
- MODEL_NAME: The model identifier to use
- HF_TOKEN: Your Hugging Face / API key
Usage:
API_BASE_URL=https://router.huggingface.co/v1 \
MODEL_NAME=Qwen/Qwen2.5-72B-Instruct \
HF_TOKEN=hf_xxx \
python inference.py
"""
import os
import sys
import json
import time
import requests
from openai import OpenAI
# ── Configuration ──
API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")
API_KEY = os.getenv("HF_TOKEN") or os.getenv("API_KEY") or os.getenv("OPENAI_API_KEY", "")
MODEL_NAME = os.getenv("MODEL_NAME", "Qwen/Qwen2.5-72B-Instruct")
# Environment URL (local or HF Space)
ENV_URL = os.getenv("ENV_URL", "http://localhost:7860")
MAX_ATTEMPTS = 3 # Max steps per task
TASKS = ["task_1", "task_2", "task_3"]
def call_environment(endpoint: str, payload: dict = None) -> dict:
"""Call the environment's HTTP API."""
url = f"{ENV_URL.rstrip('/')}/{endpoint.lstrip('/')}"
if payload is not None:
resp = requests.post(url, json=payload, timeout=30)
else:
resp = requests.get(url, timeout=30)
resp.raise_for_status()
return resp.json()
def generate_sql(client: OpenAI, schema: str, question: str, history: str = "") -> str:
"""Ask the LLM to generate a SQL query."""
system_prompt = (
"You are an expert SQL query writer. Given a database schema and a "
"natural language question, write a single SQLite-compatible SQL query "
"that answers the question.\n\n"
"RULES:\n"
"- Return ONLY the SQL query, nothing else.\n"
"- Do NOT include markdown code fences, explanations, or comments.\n"
"- Use proper SQLite syntax.\n"
"- Pay attention to column names and table relationships.\n"
"- When asked to round values, use ROUND(value, decimal_places).\n"
"- Use single quotes for string literals.\n"
)
user_prompt = f"DATABASE SCHEMA:\n{schema}\n\nQUESTION: {question}"
if history:
user_prompt += f"\n\nPREVIOUS ATTEMPTS AND FEEDBACK:\n{history}"
user_prompt += "\n\nPlease fix your query based on the feedback above."
try:
completion = client.chat.completions.create(
model=MODEL_NAME,
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt},
],
temperature=0.1,
max_tokens=500,
)
response = completion.choices[0].message.content or ""
# Clean up: remove markdown fences if present
response = response.strip()
if response.startswith("```sql"):
response = response[6:]
if response.startswith("```"):
response = response[3:]
if response.endswith("```"):
response = response[:-3]
return response.strip()
except Exception as e:
print(f" LLM API error: {e}")
return "SELECT 1"
def run_task(client: OpenAI, task_id: str) -> float:
"""Run the agent on a single task. Returns the best reward achieved."""
print(f"\n{'='*60}")
print(f"TASK: {task_id.upper()}")
print(f"{'='*60}")
# Reset environment for this task
reset_resp = call_environment("/reset", {"task_id": task_id})
obs = reset_resp["observation"]
print(f"Question: {obs['task_description']}")
print(f"Difficulty: {obs['difficulty']}")
best_reward = 0.0
history = ""
for attempt in range(1, MAX_ATTEMPTS + 1):
print(f"\n--- Attempt {attempt}/{MAX_ATTEMPTS} ---")
# Generate SQL query
sql = generate_sql(
client,
schema=obs["schema_description"],
question=obs["task_description"],
history=history,
)
print(f"SQL: {sql[:200]}{'...' if len(sql) > 200 else ''}")
# Submit to environment
step_resp = call_environment("/step", {
"task_id": task_id,
"sql_query": sql,
})
obs = step_resp["observation"]
reward = step_resp["reward"]
done = step_resp["done"]
print(f"Reward: {reward:.2f}")
print(f"Feedback: {obs['feedback']}")
best_reward = max(best_reward, reward)
# Track history for retry
history += f"\nAttempt {attempt}: SQL: {sql}\n"
history += f" Reward: {reward}, Feedback: {obs['feedback']}\n"
if obs.get("query_error"):
history += f" Error: {obs['query_error']}\n"
if done:
break
print(f"\nBest reward for {task_id}: {best_reward:.2f}")
return best_reward
def main():
"""Run the baseline agent on all tasks and report scores."""
print("=" * 60)
print("SQL Query Environment - Baseline Inference")
print("=" * 60)
print(f"API URL: {API_BASE_URL}")
print(f"Model: {MODEL_NAME}")
print(f"Env URL: {ENV_URL}")
print()
# Verify environment is up
try:
health = call_environment("/health")
print(f"Environment health: {health}")
except Exception as e:
print(f"ERROR: Cannot reach environment at {ENV_URL}: {e}")
print("Make sure the environment server is running.")
sys.exit(1)
# Create OpenAI client
client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)
# Run all tasks
results = {}
total_score = 0.0
for task_id in TASKS:
score = run_task(client, task_id)
results[task_id] = score
total_score += score
# Summary
print("\n" + "=" * 60)
print("FINAL RESULTS")
print("=" * 60)
for task_id, score in results.items():
print(f" {task_id:8s}: {score:.2f}")
avg_score = total_score / len(TASKS)
print(f" {'AVERAGE':8s}: {avg_score:.2f}")
print("=" * 60)
return results
if __name__ == "__main__":
main()