ritvik360/nl2sql-bench
0
1import os2import sys3import json4import torch5import hashlib6from pathlib import Path7from tqdm import tqdm8from transformers import AutoModelForCausalLM, AutoTokenizer9import sys10 11# --- PATCH FOR TRANSFORMERS VERSION MISMATCH ---12try:13 import transformers.activations14 if not hasattr(transformers.activations, "PytorchGELUTanh"):15 # Mapping the old name to the new existing one16 transformers.activations.PytorchGELUTanh = transformers.activations.GELUActivation17except ImportError:18 pass19# ------------------------------------------------------20 21import os22import json23import torch24# ... baaki ke saare purane imports25 26# Force script to use only the 2 free GPUs (e.g., 0 and 7)27os.environ["CUDA_VISIBLE_DEVICES"] = "0,7"28 29PROJECT_ROOT = os.path.abspath(os.path.dirname(__file__))30if PROJECT_ROOT not in sys.path:31 sys.path.insert(0, PROJECT_ROOT)32 33from data_factory.schemas import SCHEMA_CONTEXT34 35# AWQ model is 4x smaller and much faster36MODEL_NAME = "Qwen/Qwen2.5-72B-Instruct-AWQ"37INPUT_FILE = "llm_hybrid_templates.json"38OUTPUT_FILE = "nl2sql_50k_elite_dataset.jsonl"39VARIATIONS_PER_SQL = 2040BATCH_SIZE = 64 # AWQ allows much larger batches!41 42SYSTEM_PROMPT = "You are an expert SQL analyst. Write a single SELECT query that answers the question. Output ONLY the SQL query โ no markdown, no explanation, no backticks."43 44EXPANSION_PROMPT = """45You are an expert linguist and NL2SQL data augmentor. I have a SQLite database schema and a complex SQL query.46Generate exactly {count} completely different natural language questions that this exact SQL query answers.47 48RULES:49- Personas: Executive (direct), Non-tech (wordy), Analyst (technical), Curious (investigative).50- Structure: Completely change sentence flow.51- No direct column/table names.52 53DATABASE SCHEMA:54{schema}55 56SQL QUERY:57{sql}58 59OUTPUT FORMAT:60Return ONLY a valid JSON array of objects: [{{"persona": "...", "question": "..."}}]61"""62 63def extract_json_array(raw_text):64 text = raw_text.strip()65 start = text.find("[")66 end = text.rfind("]")67 if start != -1 and end != -1:68 return text[start:end+1]69 return "[]"70 71def get_hash(text):72 return hashlib.md5(text.lower().strip().encode('utf-8')).hexdigest()73 74def main():75 if not os.path.exists(INPUT_FILE):76 print(f"Error: {INPUT_FILE} not found.")77 sys.exit(1)78 79 with open(INPUT_FILE, "r") as f:80 base_templates = json.load(f)81 82 print(f"๐ Loading {MODEL_NAME} on 2 GPUs...")83 84 tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME, padding_side="left")85 tokenizer.pad_token = tokenizer.eos_token86 87 # Model loading (AWQ version automatically handles quantization)88 model = AutoModelForCausalLM.from_pretrained(89 MODEL_NAME,90 device_map="auto",91 torch_dtype=torch.float16, # AWQ models use float16/bfloat16 for weights92 low_cpu_mem_usage=True93 )94 95 seen_hashes = set()96 total_saved = 097 if os.path.exists(OUTPUT_FILE):98 with open(OUTPUT_FILE, "r") as f:99 for line in f:100 total_saved += 1 # Quick count101 102 pbar = tqdm(total=len(base_templates) * VARIATIONS_PER_SQL, initial=total_saved)103 104 # Batch processing105 for i in range(0, len(base_templates), BATCH_SIZE):106 batch = base_templates[i:i + BATCH_SIZE]107 prompts = []108 109 for temp in batch:110 msg = [111 {"role": "system", "content": "You output only JSON arrays."},112 {"role": "user", "content": EXPANSION_PROMPT.format(count=VARIATIONS_PER_SQL, schema=SCHEMA_CONTEXT[temp['domain']], sql=temp['sql'])}113 ]114 prompts.append(tokenizer.apply_chat_template(msg, tokenize=False, add_generation_prompt=True))115 116 inputs = tokenizer(prompts, return_tensors="pt", padding=True).to(model.device)117 118 try:119 with torch.no_grad():120 # Increased speed: AWQ handles large batches efficiently121 outputs = model.generate(122 **inputs,123 max_new_tokens=2048,124 temperature=0.5,125 do_sample=True,126 pad_token_id=tokenizer.eos_token_id127 )128 129 responses = tokenizer.batch_decode(outputs[:, inputs.input_ids.shape[1]:], skip_special_tokens=True)130 131 with open(OUTPUT_FILE, "a", encoding="utf-8") as out_file:132 for idx, resp in enumerate(responses):133 questions_data = json.loads(extract_json_array(resp))134 sql = batch[idx]["sql"]135 domain = batch[idx]["domain"]136 137 for item in questions_data:138 q = item.get("question", "")139 if len(q) > 10:140 q_hash = get_hash(q + sql)141 if q_hash not in seen_hashes:142 seen_hashes.add(q_hash)143 record = {144 "prompt": [145 {"role": "system", "content": SYSTEM_PROMPT},146 {"role": "user", "content": f"SCHEMA: {SCHEMA_CONTEXT[domain]}\nQUESTION: {q}"}147 ],148 "sql": sql149 }150 out_file.write(json.dumps(record, ensure_ascii=False) + "\n")151 total_saved += 1152 pbar.update(1)153 out_file.flush()154 except Exception as e:155 print(f"Batch failed: {e}")156 continue157 158 pbar.close()159 160if __name__ == "__main__":161 main()