#!/usr/bin/env python3
import os
import torch
import json
from pathlib import Path
from transformers import AutoModelForCausalLM, AutoTokenizer
from premsql.executors import SQLiteExecutor
from premsql.evaluator import Text2SQLEvaluator
from premsql.datasets import BirdDataset
from tqdm import tqdm
def get_database_schema(db_path: str) -> str:
"""Extract database schema from SQLite database"""
import sqlite3
conn = sqlite3.connect(db_path)
cursor = conn.cursor()
cursor.execute("SELECT name FROM sqlite_master WHERE type='table';")
tables = cursor.fetchall()
schema_info = []
for table in tables:
table_name = table[0]
# Quote table name to handle special characters
cursor.execute(f'PRAGMA table_info("{table_name}");')
columns = cursor.fetchall()
column_info = [f"{col[1]} ({col[2]})" for col in columns]
schema_info.append(f"Table: {table_name}\nColumns: {', '.join(column_info)}\n")
conn.close()
return "\n".join(schema_info)
def extract_sql_query(text: str) -> str:
"""Extract SQL query from model output"""
try:
sql = text.split("")[-1]
sql = sql.split("")[0]
return sql.strip()
except IndexError:
return ""
def main():
# Paths and configs
checkpoint_path = "outputs/Prem-1B-SQL-GRPO-v2/checkpoint-295"
dataset_folder = "./data"
experiment_path = Path("evaluation_results")
experiment_path.mkdir(exist_ok=True)
print("Loading model and tokenizer...")
model = AutoModelForCausalLM.from_pretrained(checkpoint_path, torch_dtype=torch.bfloat16).cuda()
tokenizer = AutoTokenizer.from_pretrained(checkpoint_path)
print("Loading evaluation dataset...")
dataset = BirdDataset(
split="validation",
dataset_folder=dataset_folder,
force_download=False
)
# Setup executor and evaluator
executor = SQLiteExecutor()
evaluator = Text2SQLEvaluator(
executor=executor,
experiment_path=experiment_path
)
# Format responses for evaluator
model_responses = []
print("\nChecking dataset structure:")
print(f"Dataset length: {len(dataset.dataset)}")
print(f"First example: {dataset.dataset[0]}")
print("\nGenerating responses...")
for example in tqdm(dataset.dataset):
# Use correct path structure for validation set
db_path = f"{dataset_folder}/bird/validation/dev_databases/{example['db_id']}/{example['db_id']}.sqlite"
if not os.path.exists(db_path):
print(f"Database not found: {db_path}")
continue
schema = get_database_schema(db_path)
prompt = f"""You are an expert SQL analyst. For the given question and database schema, generate a SQL query.
First analyze the schema and requirements, then write the query.
The database schema and tables are:
{schema}
Database: {example['db_id']}
Question: {example['question']}
Respond in the following format:
[Step by step analysis of the requirements and schema]
[Your SQL query]
"""
# Generate response
inputs = tokenizer(prompt, return_tensors="pt").to("cuda")
outputs = model.generate(
**inputs,
max_new_tokens=512,
do_sample=True,
temperature=0.7,
num_return_sequences=1,
pad_token_id=tokenizer.eos_token_id
)
response = tokenizer.decode(outputs[0], skip_special_tokens=True)
# Extract SQL query from response
sql_query = extract_sql_query(response)
if not sql_query:
sql_query = "SELECT 1" # Fallback for failed extractions
# Format response for evaluator
model_responses.append({
'question_id': len(model_responses),
'db_id': example['db_id'],
'question': example['question'],
'SQL': example['SQL'], # Ground truth
'db_path': db_path,
'generated': sql_query, # Just the SQL query
'difficulty': example.get('difficulty', 'unknown'),
'full_response': response # Keep full response for analysis
})
if not model_responses:
print("\nNo valid responses generated! Check database paths and permissions.")
return
print(f"\nGenerated {len(model_responses)} valid responses")
# Save responses
with open(experiment_path / "predict.json", "w") as f:
json.dump(model_responses, f, indent=2)
# Evaluate only if we have responses
if model_responses:
print("\nEvaluating execution accuracy...")
accuracy_results = evaluator.execute(
metric_name="accuracy",
model_responses=model_responses,
filter_by="difficulty",
meta_time_out=10,
debug=True # Add debug flag to see what's happening
)
# Convert accuracy to percentage
accuracy_formatted = {
k: v * 100 if isinstance(v, float) else v
for k, v in accuracy_results.items()
}
print(f"\nExecution Accuracy (%): {accuracy_formatted}")
print("\nEvaluating VES...")
ves_results = evaluator.execute(
metric_name="ves",
model_responses=model_responses,
filter_by="db_id",
meta_time_out=10
)
print(f"\nValid Efficiency Score: {ves_results}")
if __name__ == "__main__":
main()