ritvik360/nl2sql-bench
0
1import os
2# CRITICAL: Ye line sabse upar honi chahiye kisi bhi PyTorch import se pehle!
3os.environ["CUDA_VISIBLE_DEVICES"] = "0,7"
4
5import sys
6import torch
7from datasets import Dataset
8from transformers import AutoModelForCausalLM, AutoTokenizer
9from peft import LoraConfig
10from trl import GRPOConfig, GRPOTrainer
11
12sys.path.insert(0, "./server")
13from environment import NL2SQLEnvironment
14from models import NL2SQLAction
15from tasks import all_task_names, get_task
16
17MODEL_NAME = "Qwen/Qwen2.5-Coder-7B-Instruct"
18OUTPUT_DIR = "./qwen-7b-coder-nl2sql-grpo"
19
20SYSTEM_PROMPT = """You are a Senior Database Architect and an expert in SQLite.
21Your task is to translate natural language questions into highly optimized, correct SQLite SELECT queries.
22
23STRICT RULES:
241. Output EXACTLY ONE valid SQLite query.
252. DO NOT wrap the query in markdown formatting (no ```sql or ```).
263. DO NOT output any explanations, conversational text, or preambles (e.g., never say "Here is the query").
274. ONLY use standard SQLite functions. Avoid SQL Server, MySQL, or PostgreSQL specific syntax.
285. If the question implies ordering, use the correct ORDER BY clause.
29
30Your output must be executable directly against the database as-is."""
31
32def build_dataset():
33 data = []
34 for t_name in all_task_names():
35 task = get_task(t_name)
36 schema = task.schema_context()
37 for ex in task.examples:
38 user_content = f"SCHEMA:\n{schema}\n\nQUESTION: {ex.question}"
39 data.append({
40 "prompt": [
41 {"role": "system", "content": SYSTEM_PROMPT},
42 {"role": "user", "content": user_content}
43 ],
44 "task_name": t_name
45 })
46 return Dataset.from_list(data)
47
48def sql_reward_func(prompts, completions, task_name, **kwargs):
49 rewards = []
50 env = NL2SQLEnvironment()
51
52 for idx, completion in enumerate(completions):
53 generated_text = completion[0]['content'] if isinstance(completion, list) else completion
54
55 if generated_text.startswith("```"):
56 lines = generated_text.split("\n")
57 generated_text = "\n".join(l for l in lines if not l.strip().startswith("```")).strip()
58
59 current_task = task_name[idx] if isinstance(task_name, list) else task_name
60
61 env.reset(task_name=current_task)
62
63 try:
64 action = NL2SQLAction(query=generated_text)
65 obs = env.step(action)
66 rewards.append(float(obs.reward))
67 except Exception:
68 rewards.append(0.0)
69
70 return rewards
71
72def main():
73 dataset = build_dataset()
74
75 tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME, padding_side="right")
76 if tokenizer.pad_token is None:
77 tokenizer.pad_token = tokenizer.eos_token
78
79 model = AutoModelForCausalLM.from_pretrained(
80 MODEL_NAME,
81 torch_dtype=torch.bfloat16,
82 attn_implementation="sdpa" # Defaulting to sdpa to avoid any flash_attn setup issues
83 )
84
85 peft_config = LoraConfig(
86 r=128,
87 lora_alpha=256,
88 target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
89 bias="none",
90 task_type="CAUSAL_LM"
91 )
92
93 training_args = GRPOConfig(
94 output_dir=OUTPUT_DIR,
95 learning_rate=2e-5,
96 per_device_train_batch_size=2,
97 gradient_accumulation_steps=4,
98 max_completion_length=256,
99 num_generations=8,
100 temperature=0.5,
101 bf16=True,
102 logging_steps=5,
103 num_train_epochs=10,
104 report_to="none",
105 remove_unused_columns=False,
106 ddp_find_unused_parameters=False
107 )
108
109 trainer = GRPOTrainer(
110 model=model,
111 reward_funcs=sql_reward_func,
112 args=training_args,
113 train_dataset=dataset,
114 peft_config=peft_config,
115 processing_class=tokenizer
116 )
117
118 trainer.train()
119
120 if trainer.accelerator.is_main_process:
121 trainer.model.save_pretrained(f"{OUTPUT_DIR}/final")
122 tokenizer.save_pretrained(f"{OUTPUT_DIR}/final")
123
124if __name__ == "__main__":
125 main()