Team Ai
Apppublic

ritvik360/nl2sql-bench

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