ritvik360/nl2sql-bench
0
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()