CoolFace
Apppublic

ritvik360/nl2sql-bench

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
local_test.py118 linesDownload Raw Back to root
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())