ritvik360/nl2sql-bench
0
1import json2import os3import sys4import re5from tqdm import tqdm6 7PROJECT_ROOT = os.path.abspath(os.path.dirname(__file__))8if PROJECT_ROOT not in sys.path:9 sys.path.insert(0, PROJECT_ROOT)10 11from data_factory.validator import SQLValidator12 13INPUT_FILE = "nl2sql_50k_elite_dataset.jsonl"14OUTPUT_FILE = "nl2sql_cleaned_ready_to_train.jsonl"15 16def main():17 if not os.path.exists(INPUT_FILE):18 print(f"Error: {INPUT_FILE} not found!")19 return20 21 print(f"Sweeping dataset to remove bad SQLs...")22 23 with open(INPUT_FILE, "r", encoding="utf-8") as f:24 lines = f.readlines()25 26 validators = {}27 cleaned_count = 028 failed_count = 029 30 with open(OUTPUT_FILE, "w", encoding="utf-8") as out_f:31 for line in tqdm(lines, desc="Filtering Garbage"):32 try:33 record = json.loads(line)34 except json.JSONDecodeError:35 failed_count += 136 continue37 38 sql = record.get("sql", "").strip()39 metadata = record.get("metadata", {})40 domain = metadata.get("domain")41 42 # Fallback for domain extraction43 if not domain or domain == "unknown":44 content = record.get("prompt", [{}, {}])[1].get("content", "")45 match = re.search(r"Database:\s*([a-zA-Z0-9_]+)", content)46 domain = match.group(1) if match else "unknown"47 48 if domain == "unknown":49 failed_count += 150 continue51 52 if domain not in validators:53 validators[domain] = SQLValidator(domain, seed=42)54 55 try:56 val_result = validators[domain].validate(sql)57 # Keep ONLY if SQL is 100% perfect and returns data58 if val_result.passed and val_result.row_count > 0:59 out_f.write(line)60 cleaned_count += 161 else:62 failed_count += 163 except Exception:64 failed_count += 165 66 for v in validators.values():67 v.close()68 69 print("\n" + "="*50)70 print("DATASET CLEANUP COMPLETE")71 print("="*50)72 print(f"Original Rows : {len(lines)}")73 print(f"Cleaned Rows : {cleaned_count} (100% Valid SQL)")74 print(f"Removed Rows : {failed_count}")75 print(f"Saved To : {OUTPUT_FILE}")76 print("="*50)77 78if __name__ == "__main__":79 main()