ritvik360/nl2sql-bench
0
1import asyncio
2import os
3import sys
4import torch
5from transformers import AutoModelForCausalLM, AutoTokenizer
6from peft import PeftModel
7
8# --- Configuration ---
9BASE_MODEL = "Qwen/Qwen2.5-0.5B-Instruct"
10LORA_DIR = "./qwen-nl2sql-grpo/checkpoint-50"
11SPACE_URL = "http://localhost:8000" # Local server URL
12TASKS = ["simple-filter", "join-aggregation", "analytics-window"]
13MAX_STEPS = 5
14
15print("Loading Base Model and LoRA weights...")
16tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL)
17base_model = AutoModelForCausalLM.from_pretrained(
18 BASE_MODEL,
19 torch_dtype=torch.bfloat16,
20 device_map="auto"
21)
22model = PeftModel.from_pretrained(base_model, LORA_DIR)
23
24# --- System Prompt & LLM Call ---
25SYSTEM_PROMPT = """You are an expert SQL analyst working with a SQLite e-commerce database.
26Write a single SELECT query. Output ONLY the SQL query, nothing else. No markdown."""
27
28def call_local_llm(user_prompt: str) -> str:
29 messages = [
30 {"role": "system", "content": SYSTEM_PROMPT},
31 {"role": "user", "content": user_prompt}
32 ]
33 text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
34 inputs = tokenizer([text], return_tensors="pt").to(model.device)
35
36 with torch.no_grad():
37 outputs = model.generate(**inputs, max_new_tokens=256, temperature=0.2, do_sample=True)
38
39 response = tokenizer.decode(outputs[0][inputs.input_ids.shape[1]:], skip_special_tokens=True)
40
41 # Strip markdown code fences if model wraps in ```sql ... ```
42 if response.startswith("```"):
43 lines = response.split("\n")
44 response = "\n".join(l for l in lines if not l.strip().startswith("```")).strip()
45 return response if response else "SELECT 1"
46
47def build_user_prompt(question, schema_context, step, last_query, last_error, last_result, result_columns):
48 parts = [f"QUESTION: {question}", ""]
49 if step > 1:
50 parts.append(f"Your previous SQL (step {step - 1}):")
51 parts.append(f" {' '.join(last_query.split())}")
52 parts.append("")
53 if last_error:
54 parts.append(f"ERROR: {last_error}")
55 elif last_result:
56 preview = str(last_result[:3]).replace("\n", " ")
57 parts.append(f"RESULT PREVIEW (first 3 rows): {preview}")
58 parts.append(f"COLUMNS: {result_columns}")
59 parts.append("")
60 parts.append("Please correct or refine your query.")
61 else:
62 parts.append("Write a SQL query to answer the question.")
63 return "\n".join(parts)
64
65async def main():
66 from client import NL2SQLEnv, NL2SQLAction
67
68 all_results = []
69
70 for task_name in TASKS:
71 print(f"\n--- Starting Task: {task_name} ---")
72 os.environ["NL2SQL_DEFAULT_TASK"] = task_name
73
74 try:
75 async with NL2SQLEnv(base_url=SPACE_URL) as env:
76 result = await env.reset()
77 obs = result.observation
78
79 rewards = []
80 success = False
81
82 for step in range(1, MAX_STEPS + 1):
83 if obs.done:
84 break
85
86 user_prompt = build_user_prompt(
87 obs.question, obs.schema_context, step,
88 obs.last_query, obs.last_error, obs.last_result, obs.result_columns
89 )
90
91 sql = call_local_llm(user_prompt)
92
93 print(f"Step {step} Agent Output: {sql}")
94
95 step_result = await env.step(NL2SQLAction(query=sql))
96 obs = step_result.observation
97
98 reward = obs.reward or 0.0
99 rewards.append(reward)
100 print(f"Step {step} Reward: {reward}")
101
102 if obs.done:
103 break
104
105 score = sum(rewards) / max(len(rewards), 1)
106 success = score >= 0.7
107 print(f"Final Score for {task_name}: {score:.3f}")
108 all_results.append({"task": task_name, "score": score, "success": success})
109
110 except Exception as e:
111 print(f"Error testing task {task_name}: {e}")
112
113 print("\n=== Final Results ===")
114 for r in all_results:
115 print(f"{r['task']}: Score {r['score']:.3f} | Success: {r['success']}")
116
117if __name__ == "__main__":
118 asyncio.run(main())