Team Ai
Apppublic

ritvik360/nl2sql-bench

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