CoolFace
Apppublic

ritvik360/nl2sql-bench

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