Team Ai
Apppublic

ritvik360/nl2sql-bench

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
custom_train.py250 linesDownload Raw Back to root
1"""2merge_and_train.py3==================41. Merges nl2sql_cleaned_ready_to_train.jsonl + edge_cases.jsonl52. Shuffles the combined dataset63. Retrains using the same GRPO setup as train.py7 8Run:9    python merge_and_train.py10 11Flags (env vars):12    EDGE_FILE   — path to edge cases jsonl   (default: edge_cases.jsonl)13    BASE_FILE   — path to existing cleaned    (default: nl2sql_cleaned_ready_to_train.jsonl)14    MERGED_FILE — merged output path          (default: nl2sql_merged_final.jsonl)15    SKIP_MERGE  — set "1" to skip merge step and go straight to training16"""17 18import os, sys, json, random19import torch20from datasets import Dataset21from transformers import AutoModelForCausalLM, AutoTokenizer22from peft import LoraConfig23from trl import GRPOConfig, GRPOTrainer24 25os.environ["CUDA_VISIBLE_DEVICES"] = "0,5,1,6"26 27sys.path.insert(0, "./server")28from environment import NL2SQLEnvironment29from models import NL2SQLAction30from tasks import all_task_names, get_task31 32# ── Config ───────────────────────────────────────────────────────────────────33BASE_FILE   = os.getenv("BASE_FILE",   "nl2sql_cleaned_ready_to_train.jsonl")34EDGE_FILE   = os.getenv("EDGE_FILE",   "edge_cases.jsonl")35MERGED_FILE = os.getenv("MERGED_FILE", "nl2sql_merged_final.jsonl")36SKIP_MERGE  = os.getenv("SKIP_MERGE",  "0") == "1"37 38MODEL_NAME  = "Qwen/Qwen2.5-Coder-7B-Instruct"39OUTPUT_DIR  = "./qwen-7b-coder-nl2sql-grpo-v2"40 41SYSTEM_PROMPT = """You are a Senior Database Architect and an expert in SQLite.42Your task is to translate natural language questions into highly optimized, correct SQLite SELECT queries.43 44STRICT RULES:451. Output EXACTLY ONE valid SQLite query.462. DO NOT wrap the query in markdown formatting (no ```sql or ```).473. DO NOT output any explanations, conversational text, or preambles.484. ONLY use standard SQLite functions.495. If the question implies ordering, use the correct ORDER BY clause.506. SELECT only the columns explicitly requested — no extras.51 52Your output must be executable directly against the database as-is."""53 54 55# ── Step 1: Merge ─────────────────────────────────────────────────────────────56 57def merge_datasets():58    if SKIP_MERGE:59        print(f"[SKIP_MERGE=1] Using existing {MERGED_FILE}")60        return61 62    print(f"Loading base:  {BASE_FILE}")63    print(f"Loading edges: {EDGE_FILE}")64 65    base_lines = []66    with open(BASE_FILE, "r", encoding="utf-8") as f:67        for line in f:68            line = line.strip()69            if line:70                base_lines.append(line)71 72    edge_lines = []73    with open(EDGE_FILE, "r", encoding="utf-8") as f:74        for line in f:75            line = line.strip()76            if line:77                edge_lines.append(line)78 79    combined = base_lines + edge_lines80    random.shuffle(combined)81 82    with open(MERGED_FILE, "w", encoding="utf-8") as f:83        for line in combined:84            f.write(line + "\n")85 86    print(87        f"Merged: {len(base_lines)} base + {len(edge_lines)} edge "88        f"= {len(combined)} total → {MERGED_FILE}"89    )90 91 92# ── Step 2: Build HF Dataset ──────────────────────────────────────────────────93 94def build_dataset():95    """96    Primary source: merged JSONL (base + edge cases).97    Fallback: task examples from server/tasks/ (same as original train.py).98    Both are combined so GRPO sees everything.99    """100    data = []101 102    # Load merged JSONL103    with open(MERGED_FILE, "r", encoding="utf-8") as f:104        for line in f:105            line = line.strip()106            if not line:107                continue108            rec = json.loads(line)109            # rec has "prompt" (list of messages) and "sql"110            # GRPO needs "prompt" and "task_name" — we use a synthetic task_name111            data.append({112                "prompt":    rec["prompt"],113                "task_name": "merged_jsonl"   # grader falls back to execution-based reward114            })115 116    # Also keep the original task examples so GRPO reward env works for them117    for t_name in all_task_names():118        task  = get_task(t_name)119        schema = task.schema_context()120        for ex in task.examples:121            data.append({122                "prompt": [123                    {"role": "system", "content": SYSTEM_PROMPT},124                    {"role": "user",   "content": f"SCHEMA:\n{schema}\n\nQUESTION: {ex.question}"}125                ],126                "task_name": t_name127            })128 129    random.shuffle(data)130    print(f"Dataset size: {len(data)} samples")131    return Dataset.from_list(data)132 133 134# ── Step 3: Reward function ───────────────────────────────────────────────────135 136def sql_reward_func(prompts, completions, task_name, **kwargs):137    rewards = []138    env     = NL2SQLEnvironment()139 140    for idx, completion in enumerate(completions):141        generated = (142            completion[0]["content"] if isinstance(completion, list) else completion143        )144        # Strip code fences defensively145        import re146        generated = re.sub(r"```(?:sql)?\n?(.*?)```", r"\1", generated, flags=re.DOTALL).strip()147 148        t = task_name[idx] if isinstance(task_name, list) else task_name149 150        # For merged_jsonl rows the env won't have a matching task →151        # reward purely on execution (non-empty result set = +1, error = 0)152        if t == "merged_jsonl":153            rewards.append(_execution_reward(generated, prompts[idx]))154            continue155 156        env.reset(task_name=t)157        try:158            obs = env.step(NL2SQLAction(query=generated))159            rewards.append(float(obs.reward))160        except Exception:161            rewards.append(0.0)162 163    return rewards164 165 166def _execution_reward(sql: str, prompt) -> float:167    """Simple execution check for merged_jsonl samples."""168    import sqlite3, re as _re169 170    # Extract schema from the user message171    user_content = ""172    for msg in (prompt if isinstance(prompt, list) else []):173        if isinstance(msg, dict) and msg.get("role") == "user":174            user_content = msg.get("content", "")175            break176 177    schema_match = _re.search(r"SCHEMA:\s*(.*?)\nQUESTION:", user_content, _re.DOTALL)178    if not schema_match:179        return 0.5  # can't verify, neutral reward180 181    schema_sql = schema_match.group(1).strip()182    try:183        conn = sqlite3.connect(":memory:")184        conn.executescript(schema_sql)185        rows = conn.execute(sql).fetchall()186        conn.close()187        return 1.0 if rows else 0.3   # ran cleanly but empty → partial credit188    except Exception:189        return 0.0190 191 192# ── Step 4: Train ─────────────────────────────────────────────────────────────193 194def main():195    merge_datasets()196    dataset = build_dataset()197 198    tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME, padding_side="right")199    if tokenizer.pad_token is None:200        tokenizer.pad_token = tokenizer.eos_token201 202    model = AutoModelForCausalLM.from_pretrained(203        MODEL_NAME,204        torch_dtype=torch.bfloat16,205        attn_implementation="sdpa"206    )207 208    peft_config = LoraConfig(209        r=128,210        lora_alpha=256,211        target_modules=["q_proj","k_proj","v_proj","o_proj","gate_proj","up_proj","down_proj"],212        bias="none",213        task_type="CAUSAL_LM"214    )215 216    training_args = GRPOConfig(217        output_dir=OUTPUT_DIR,218        learning_rate=1e-5,              # lower LR for fine-grained edge case tuning219        per_device_train_batch_size=2,220        gradient_accumulation_steps=4,221        max_completion_length=256,222        num_generations=8,223        temperature=0.5,224        bf16=True,225        logging_steps=5,226        num_train_epochs=5,              # fewer epochs — base knowledge already there227        report_to="none",228        remove_unused_columns=False,229        ddp_find_unused_parameters=False230    )231 232    trainer = GRPOTrainer(233        model=model,234        reward_funcs=sql_reward_func,235        args=training_args,236        train_dataset=dataset,237        peft_config=peft_config,238        processing_class=tokenizer239    )240 241    trainer.train()242 243    if trainer.accelerator.is_main_process:244        trainer.model.save_pretrained(f"{OUTPUT_DIR}/final")245        tokenizer.save_pretrained(f"{OUTPUT_DIR}/final")246        print(f"\nSaved to {OUTPUT_DIR}/final")247 248 249if __name__ == "__main__":250    main()