from unsloth import FastLanguageModel, PatchFastRL PatchFastRL("GRPO", FastLanguageModel) from unsloth import is_bfloat16_supported import torch #---v4 #pip install diffusers #!pip install "unsloth==2025.2.4" vllm #!pip install --upgrade pillow num_generations = 8 per_device_train_batch_size = 8 gradient_accumulation_steps = 4 num_train_epochs = 3 beta = 0.01 learning_rate = 1e-4 max_seq_length = 1650 lora_rank = 8 lora_alpha = lora_rank * 2 dataset_dir = "cot-qa-distill.csv" model_name = "./models/deepquery4-sft-1" output_dir = "./result/deepquery4-grpo-1" SYS="""You are DeepQuery, a data science expert. Below, you are presented with a database schema, a question and a hint.Your task is to read the schema with annotations of the columns, understand the question and the hint, and generate a valid SQL query to answer the question. You should reason step by step, and includes your reasonings between and .""" model, tokenizer = FastLanguageModel.from_pretrained( model_name = model_name, max_seq_length = max_seq_length, load_in_4bit = True, # False for LoRA 16bit fast_inference = True, # Enable vLLM fast inference max_lora_rank = lora_rank, gpu_memory_utilization = 0.4, # Reduce if out of memory ) tokenizer.pad_token = tokenizer.eos_token model = FastLanguageModel.get_peft_model( model, r = lora_rank, # Choose any number > 0 ! Suggested 8, 16, 32, 64, 128 target_modules = [ "gate_proj","up_proj", "down_proj","q_proj", "k_proj","o_proj", "v_proj" #"q_proj", "k_proj", "up_proj", "down_proj", ], # Remove QKVO if out of memory lora_alpha = lora_alpha, use_gradient_checkpointing = "unsloth", # Enable long context finetuning random_state = 3407, ) def extract_answer(text: str) -> str: start_tag = 'My final answer is: \n```sql\n' end_tag = '\n```' start_tag_index = text.find(start_tag) if start_tag_index != -1: start_index = start_tag_index + len(start_tag) end_index = text.find(end_tag, start_index) if end_index != -1: return text[start_index:end_index].strip() return "" from typing import List, Dict def sqlparser(left, right): def format_sql(sql): keywords = ['SELECT', 'FROM', 'WHERE', 'GROUP BY', 'ORDER BY', 'JOIN', 'INNER JOIN', 'HAVING'] formatted_sql = sql.upper() for keyword in keywords: formatted_sql = re.sub(rf'\b{keyword}\b', keyword, formatted_sql, flags=re.IGNORECASE) return formatted_sql left = format_sql(left) right = format_sql(right) components = { 'SELECT': r'SELECT\s+(.+?)\s+FROM', 'FROM': r'FROM\s+(.+?)(?=\s+(?:WHERE|GROUP BY|ORDER BY|JOIN|INNER JOIN|HAVING|$))', 'JOIN': r'(?:INNER )?JOIN\s+(.+?)\s+ON\s+(.+?)(?=\s*(?:WHERE|GROUP BY|ORDER BY|JOIN|INNER JOIN|HAVING|$))', 'WHERE': r'WHERE\s+(.+?)(?=\s+(?:GROUP BY|ORDER BY|HAVING|$))', 'GROUP BY': r'GROUP BY\s+(.+?)(?=\s+(?:ORDER BY|HAVING|$))', 'HAVING': r'HAVING\s+(.+?)(?=\s+(?:ORDER BY|$))', 'ORDER BY': r'ORDER BY\s+(.+)$' } def parse_component(sql, component, pattern): if component == 'SELECT': match = re.search(pattern, sql, re.IGNORECASE) if match: elements = match.group(1).split(',') return set(element.strip() for element in elements) return set() elif component == 'JOIN': joins = [] for match in re.finditer( r'(?:INNER )?JOIN\s+(.+?)\s+ON\s+(.+?)(?=\s*(?:WHERE|GROUP BY|ORDER BY|JOIN|INNER JOIN|HAVING|$))', sql, re.IGNORECASE ): joins.append((match.group(1).strip(), match.group(2).strip())) return joins elif component in ['GROUP BY', 'ORDER BY']: match = re.search(pattern, sql, re.IGNORECASE) if match: elements = match.group(1).split(',') return set(element.strip() for element in elements) return set() else: match = re.search(pattern, sql, re.IGNORECASE) return match.group(1).strip() if match else '' left_components = {} right_components = {} for component, pattern in components.items(): left_components[component] = parse_component(left, component, pattern) right_components[component] = parse_component(right, component, pattern) score = 0 total = 0 for component in components: lc = left_components[component] rc = right_components[component] if lc or rc: total += 1 if lc == rc: score += 1 return score / total if total != 0 else 0 def correctness_reward_func(prompts, completions, answer, **kwargs) -> list[float]: responses = [completion[0]['content'] for completion in completions] extracted_responses = [extract_answer(r) for r in responses] return [sqlparser(a, r) for r, a in zip(extracted_responses, answer)] def bonus_reward_func(prompts, completions, answer, **kwargs) -> list[float]: responses = [completion[0]['content'] for completion in completions] extracted_responses = [extract_answer(r) for r in responses] return [1.0 if r == a else 0.0 for r, a in zip(extracted_responses, answer)] def strict_format_reward_func(completions, **kwargs) -> list[float]: """Reward function that checks if the completion has a specific format and minimum content length.""" pattern = r"^(.*?).*" responses = [completion[0]["content"] for completion in completions] rewards = [] for r in responses: match = re.match(pattern, r, re.DOTALL) if match: reasoning = match.group(1) if len(reasoning.strip()) < 300: rewards.append(0.0) else: rewards.append(0.1) else: rewards.append(0.0) return rewards from trl import GRPOConfig, GRPOTrainer training_args = GRPOConfig( use_vllm = True, # use vLLM for fast inference learning_rate = learning_rate, adam_beta1 = 0.9, adam_beta2 = 0.999, weight_decay = 0.1, warmup_ratio = 0.1, beta=beta, lr_scheduler_type = "cosine", optim = "paged_adamw_8bit", logging_steps = 1, bf16 = is_bfloat16_supported(), fp16 = not is_bfloat16_supported(), per_device_train_batch_size = per_device_train_batch_size, gradient_accumulation_steps = gradient_accumulation_steps, # Increase to 4 for smoother training num_generations = num_generations, # Decrease if out of memory max_prompt_length = max_seq_length, max_completion_length = max_seq_length//2, num_train_epochs = num_train_epochs, # Set to 1 for a full training run save_steps = 100, max_grad_norm = 0.5, report_to = "wandb", # Can use Weights & Biases output_dir = output_dir, ) import re from datasets import load_dataset, Dataset import pandas as pd def dataset_process(name: str = dataset_dir) -> List[Dict]: data = pd.read_csv(name) processed_data = data.apply(lambda row: { 'prompt': [ {'role': 'system', 'content': SYS}, {'role': 'user', 'content': row['query']} ], 'answer': row['answer'] }, axis=1).tolist() return processed_data dataset = dataset_process() from swanlab.integration.transformers import SwanLabCallback swanlab_callback = SwanLabCallback( project = "deepquery4-grpo", experiment_name = "deepquery4-grpo-1", ) trainer = GRPOTrainer( model = model, processing_class = tokenizer, reward_funcs = [ strict_format_reward_func, correctness_reward_func, bonus_reward_func ], args = training_args, train_dataset = dataset, callbacks = [swanlab_callback] ) trainer.train(resume_from_checkpoint = True)