File size: 4,625 Bytes
3b4b411
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
---
base_model: Qwen/Qwen3-4B-Base
library_name: peft
pipeline_tag: text-generation
language:
- en
tags:
- base_model:adapter:Qwen/Qwen3-4B-Base
- peft
- lora
- text-to-sql
- sql
- grpo
- rlvr
- trl
---

# SQL-R1-GRPO

SQL-R1-GRPO is the main SQL-R1 checkpoint: a LoRA adapter trained on
Qwen3-4B-Base with completion-only supervised fine-tuning followed by GRPO
using verifiable SQLite execution feedback.

The model receives a database schema, optional evidence, and a natural-language
question, then generates one read-only SQLite query. This repository contains
the adapter only; load it together with
[`Qwen/Qwen3-4B-Base`](https://huggingface.co/Qwen/Qwen3-4B-Base).

## Training

GRPO starts from the SQL-R1-SFT adapter and optimizes generated SQL with a
layered execution reward:

| Outcome | Reward |
|---|---:|
| Empty output | -1.00 |
| Invalid or non-read-only SQL | -0.75 |
| SQL execution error | -0.25 |
| Truncated execution result | -0.10 |
| Executable but incorrect result | +0.25 |
| Execution-equivalent to the gold query | +1.00 |

A low-weight format reward additionally encourages concise SQL-only output.
Gold SQL and gold execution results are used only by the reward function and
are never included in the model prompt.

- RL training records: 12,885
- Optimization steps: 300
- Candidates per prompt: 4
- Approximate generated rollouts: 1,200
- Learning rate: 5e-6
- GRPO beta: 0.001
- Clipping epsilon: 0.2
- Precision: BF16
- Execution backend: parallel read-only SQLite workers

## Evaluation

Execution accuracy (EX) and execution-valid rate were measured on the complete
BIRD Dev and Spider Dev splits with the same greedy decoding and execution
pipeline at every stage.

| Checkpoint | BIRD Dev EX | BIRD execution-valid | Spider Dev EX | Spider execution-valid |
|---|---:|---:|---:|---:|
| SQL-R1-SFT | 42.37% | 87.29% | 77.27% | 96.62% |
| **SQL-R1-GRPO** | **42.37%** | **87.35%** | **79.59%** | **97.00%** |

GRPO improves Spider Dev EX by 2.32 percentage points over SFT while preserving
BIRD Dev EX. This is the recommended SQL-R1 checkpoint.

These numbers are project-side development-set measurements, not hidden-test
leaderboard submissions. Results from systems using prompting, database
contents, self-consistency, reranking, or larger proprietary models are not
directly comparable.

## Usage

```python
import torch
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer

base_model_id = "Qwen/Qwen3-4B-Base"
repo_id = "j2521402/SQL-R1"
adapter_subfolder = "grpo"

tokenizer = AutoTokenizer.from_pretrained(
    repo_id,
    subfolder=adapter_subfolder,
)
model = AutoModelForCausalLM.from_pretrained(
    base_model_id,
    torch_dtype=torch.bfloat16,
    device_map="auto",
)
model = PeftModel.from_pretrained(
    model,
    repo_id,
    subfolder=adapter_subfolder,
)
model.eval()

messages = [
    {
        "role": "system",
        "content": (
            "You are a Text-to-SQL assistant. Given a SQLite database schema, "
            "optional evidence, and a question, return exactly one read-only "
            "SQLite query. Do not include explanations or Markdown fences."
        ),
    },
    {
        "role": "user",
        "content": (
            'Database schema:\nTABLE "singer" ("Singer_ID" INT, "Name" TEXT);\n\n'
            "Question:\nHow many singers do we have?"
        ),
    },
]

inputs = tokenizer.apply_chat_template(
    messages,
    add_generation_prompt=True,
    return_tensors="pt",
    return_dict=True,
).to(model.device)

with torch.inference_mode():
    output = model.generate(
        **inputs,
        max_new_tokens=256,
        do_sample=False,
        eos_token_id=tokenizer.eos_token_id,
        pad_token_id=tokenizer.pad_token_id,
    )

completion = output[0, inputs["input_ids"].shape[1]:]
print(tokenizer.decode(completion, skip_special_tokens=True).strip())
```

## Limitations

- Training and evaluation focus on English Text-to-SQL and SQLite.
- Prompts contain database schemas and optional evidence, but do not perform
  automatic database-content retrieval.
- Execution equivalence is a strong but imperfect semantic correctness signal.
- The model may still generate incorrect or expensive queries; execute outputs
  in a read-only sandbox with time and row limits.
- This is a direct Text-to-SQL model, not a general tool-calling agent.

## Project

Training, data processing, reward, and evaluation code:
[j2521402/SQL-R1](https://github.com/j2521402/SQL-R1).

Core libraries: PyTorch 2.8.0, Transformers 5.14.1, PEFT 0.19.1,
TRL 1.8.0, Datasets 5.0.0, and Accelerate 1.14.0.