#!/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()