211 lines
7.3 KiB
Markdown
211 lines
7.3 KiB
Markdown
|
|
---
|
||
|
|
library_name: transformers
|
||
|
|
pipeline_tag: text-generation
|
||
|
|
language:
|
||
|
|
- en
|
||
|
|
tags:
|
||
|
|
- text2sql
|
||
|
|
- sql
|
||
|
|
- pgvector
|
||
|
|
- grpo
|
||
|
|
- lora
|
||
|
|
- safetensors
|
||
|
|
- bf16
|
||
|
|
---
|
||
|
|
|
||
|
|
# text2sql-7b-v3-16
|
||
|
|
|
||
|
|
## Model Details
|
||
|
|
|
||
|
|
`text2sql-7b-v3-16` is an 8B-parameter causal language model fine-tuned for text-to-SQL generation. It accepts a natural-language question plus schema context and returns SQL inside the notebook's required answer format.
|
||
|
|
|
||
|
|
- Model type: causal language model, text generation
|
||
|
|
- Primary task: text-to-SQL
|
||
|
|
- Base model family: `Arctic-Text2SQL-R1-7B`
|
||
|
|
- Output format: `<think>...</think><answer>SQL</answer>`
|
||
|
|
- Tensor format: Safetensors
|
||
|
|
- Precision: `BF16`
|
||
|
|
- Fine-tuning method: PEFT LoRA adapters merged into the final model artifact
|
||
|
|
- Final artifact path used locally: `outputs/text2sql-7b-v3-16/merged-text2sql-7b-v3-16`
|
||
|
|
|
||
|
|
## Intended Technical Use
|
||
|
|
|
||
|
|
Use this model to generate PostgreSQL-style SQL from a natural-language question and schema description. The fine-tuning set contains both plain SQL examples and pgvector retrieval examples using `embed_query()` with vector-distance operators.
|
||
|
|
|
||
|
|
This model is not a SQL execution engine. Generated SQL should be parsed, reviewed, and run in a controlled environment before use.
|
||
|
|
|
||
|
|
## How to Use
|
||
|
|
|
||
|
|
```python
|
||
|
|
import torch
|
||
|
|
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||
|
|
|
||
|
|
model_id = "ihebaker10/text2sql-7b-v3-16"
|
||
|
|
|
||
|
|
tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
|
||
|
|
model = AutoModelForCausalLM.from_pretrained(
|
||
|
|
model_id,
|
||
|
|
torch_dtype=torch.bfloat16,
|
||
|
|
device_map="auto",
|
||
|
|
trust_remote_code=True,
|
||
|
|
)
|
||
|
|
|
||
|
|
messages = [
|
||
|
|
{
|
||
|
|
"role": "system",
|
||
|
|
"content": "Generate SQL. Return the final SQL inside <answer>...</answer>.",
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"role": "user",
|
||
|
|
"content": "Schema:\n...\n\nQuestion:\nList the latest 10 records.",
|
||
|
|
},
|
||
|
|
]
|
||
|
|
|
||
|
|
prompt = tokenizer.apply_chat_template(
|
||
|
|
messages,
|
||
|
|
tokenize=False,
|
||
|
|
add_generation_prompt=True,
|
||
|
|
)
|
||
|
|
inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
|
||
|
|
|
||
|
|
with torch.no_grad():
|
||
|
|
outputs = model.generate(
|
||
|
|
**inputs,
|
||
|
|
max_new_tokens=768,
|
||
|
|
temperature=0.2,
|
||
|
|
top_p=0.95,
|
||
|
|
do_sample=False,
|
||
|
|
)
|
||
|
|
|
||
|
|
print(tokenizer.decode(outputs[0], skip_special_tokens=True))
|
||
|
|
```
|
||
|
|
|
||
|
|
## Training Data
|
||
|
|
|
||
|
|
The training data contains 10,000 text-to-SQL examples loaded from a MongoDB collection and split into 9,001 training examples and 999 validation examples.
|
||
|
|
|
||
|
|
| Split / Type | Examples | Share |
|
||
|
|
|---|---:|---:|
|
||
|
|
| Train | 9,001 | 90.0% |
|
||
|
|
| Validation | 999 | 10.0% |
|
||
|
|
| SQL-only | 3,931 | 39.3% |
|
||
|
|
| pgvector | 6,069 | 60.7% |
|
||
|
|
|
||
|
|
Each record contains an identifier, `is_sql_only`, natural-language question, reference SQL query, reasoning hint, schema text, table list, and version.
|
||
|
|
|
||
|
|
## Training Procedure
|
||
|
|
|
||
|
|
Training used a three-stage PEFT workflow:
|
||
|
|
|
||
|
|
| Stage | Objective | Data | Schedule |
|
||
|
|
|---|---|---|---|
|
||
|
|
| Stage 1a | SQL-only SFT warm-up | SQL-only examples | 1 epoch |
|
||
|
|
| Stage 1b | Mixed SFT | Natural mixed SQL + pgvector split | 2 epochs |
|
||
|
|
| Stage 2 | GRPO refinement | Natural mixed SQL + pgvector split | about 0.5 epoch, capped by max steps |
|
||
|
|
|
||
|
|
The LoRA adapter was merged into the base weights before publishing the final Safetensors artifact.
|
||
|
|
|
||
|
|
## Hyperparameters
|
||
|
|
|
||
|
|
| Parameter | Stage 1a SFT | Stage 1b SFT | Stage 2 GRPO |
|
||
|
|
|---|---:|---:|---:|
|
||
|
|
| Learning rate | `1e-4` | `7e-5` | `1e-05` |
|
||
|
|
| Scheduler | `cosine` | `cosine` | `SchedulerType.COSINE` |
|
||
|
|
| Warmup ratio | `0.05` | `0.03` | `0.05` |
|
||
|
|
| Per-device train batch size | `1` | `1` | `2` |
|
||
|
|
| Gradient accumulation | `16` | `16` | `4` |
|
||
|
|
| Effective batch size | `16` | `16` | `8` |
|
||
|
|
| Weight decay | `0.01` | `0.01` | `0.01` |
|
||
|
|
| Max grad norm | `1.0` | `1.0` | `1.0` |
|
||
|
|
| Precision | `bf16` | `bf16` | `bf16` |
|
||
|
|
| Optimizer | `adamw_torch` | `adamw_torch` | `adamw_torch` |
|
||
|
|
| Max sequence length | `1536` | `1536` | n/a |
|
||
|
|
| Max prompt length | n/a | n/a | `1024` |
|
||
|
|
| Max completion length | n/a | n/a | `192` |
|
||
|
|
| Rollouts per prompt | n/a | n/a | `2` |
|
||
|
|
| Temperature | n/a | n/a | `0.8` |
|
||
|
|
| Top-p | n/a | n/a | `0.95` |
|
||
|
|
| KL beta | n/a | n/a | `0.002` |
|
||
|
|
| Early stopping patience | n/a | n/a | `2` |
|
||
|
|
|
||
|
|
LoRA configuration:
|
||
|
|
|
||
|
|
| Parameter | Value |
|
||
|
|
|---|---:|
|
||
|
|
| Rank | `64` |
|
||
|
|
| Alpha | `64` |
|
||
|
|
| Dropout | `0.05` |
|
||
|
|
| Bias | `none` |
|
||
|
|
| Target modules | `down_proj`, `gate_proj`, `k_proj`, `o_proj`, `q_proj`, `up_proj`, `v_proj` |
|
||
|
|
|
||
|
|
## Reward Functions
|
||
|
|
|
||
|
|
Stage 2 used the following reward checks:
|
||
|
|
|
||
|
|
- `format_reward`: validates the `<think>` and `<answer>` response structure.
|
||
|
|
- `sql_syntax_reward`: rewards SQL that parses successfully.
|
||
|
|
- `pgvector_usage_reward`: checks whether `embed_query()` is used only when expected.
|
||
|
|
- `execution_reward`: optionally compares generated SQL results with reference SQL when `DB_URL` is configured.
|
||
|
|
|
||
|
|
## Evaluation
|
||
|
|
|
||
|
|
Positive class for precision, recall, and F1 is pgvector usage, meaning the generated SQL contains `embed_query()`.
|
||
|
|
|
||
|
|
| Split | N | Format | Syntax | Use accuracy | Precision | Recall | F1 | Exact match |
|
||
|
|
|---|---:|---:|---:|---:|---:|---:|---:|---:|
|
||
|
|
| SQL-only | 393 | 100.0% | 99.5% | 89.8% | 0.0% | 0.0% | 0.0% | 33.1% |
|
||
|
|
| pgvector | 606 | 100.0% | 99.0% | 73.8% | 100.0% | 73.8% | 84.9% | 11.7% |
|
||
|
|
| combined | 999 | 100.0% | 99.2% | 80.1% | 91.8% | 73.8% | 81.8% | 20.1% |
|
||
|
|
|
||
|
|
Stage eval losses:
|
||
|
|
|
||
|
|
| Stage | Eval loss |
|
||
|
|
|---|---:|
|
||
|
|
| Stage 1a | `0.013554` |
|
||
|
|
| Stage 1b | `0.002467` |
|
||
|
|
| Stage 2 | `0.000403` |
|
||
|
|
|
||
|
|
Detailed pgvector decision metrics from the saved final merged-model evaluation:
|
||
|
|
|
||
|
|
| Split | TP | FP | FN | TN | SQL-only specificity | Embed false positive rate |
|
||
|
|
|---|---:|---:|---:|---:|---:|---:|
|
||
|
|
| SQL-only | 0 | 40 | 0 | 353 | 89.8% | 10.2% |
|
||
|
|
| pgvector | 447 | 0 | 159 | 0 | 0.0% | 0.0% |
|
||
|
|
| combined | 447 | 40 | 159 | 353 | 89.8% | 10.2% |
|
||
|
|
|
||
|
|
Detailed saved stage metrics:
|
||
|
|
|
||
|
|
| Stage | Train loss | Eval loss | Train runtime | Eval runtime | Train samples/s | Eval samples/s | Train steps/s | Eval steps/s | Total FLOPs |
|
||
|
|
|---|---:|---:|---:|---:|---:|---:|---:|---:|---:|
|
||
|
|
| Stage 1a | `2.460437` | `0.013554` | 2478s | 167.1s | 1.428 | 5.979 | 0.089 | 5.979 | `2.222e+17` |
|
||
|
|
| Stage 1b | `0.049128` | `0.002467` | 12935s | 164.9s | 1.392 | 6.057 | 0.087 | 6.057 | `1.165e+18` |
|
||
|
|
| Stage 2 | `0.000392` | `0.000403` | 36192s | 2772.4s | 0.124 | 0.360 | 0.016 | 0.045 | `0.000e+00` |
|
||
|
|
|
||
|
|
## Limitations
|
||
|
|
|
||
|
|
- SQL should be validated before execution.
|
||
|
|
- Exact-match SQL is strict and may undercount semantically equivalent queries.
|
||
|
|
- pgvector recall on the validation split is lower than precision, so some retrieval-style prompts may be answered as plain SQL.
|
||
|
|
- Execution-based reward depends on `DB_URL`; without it, training optimizes formatting, SQL syntax, and pgvector usage only.
|
||
|
|
|
||
|
|
## Technical Environment
|
||
|
|
|
||
|
|
- Frameworks: `transformers`, `peft`, `trl`, `datasets`, `torch`
|
||
|
|
- Attention implementation: `sdpa`
|
||
|
|
- Gradient checkpointing: enabled
|
||
|
|
- CUDA memory fraction in notebook: `0.98`
|
||
|
|
- TF32 matmul: enabled
|
||
|
|
- TensorBoard logging: enabled for all training stages
|
||
|
|
|
||
|
|
## Environmental Impact
|
||
|
|
|
||
|
|
Carbon emissions were not measured. Recorded training runtimes were approximately:
|
||
|
|
|
||
|
|
| Stage | Runtime |
|
||
|
|
|---|---:|
|
||
|
|
| Stage 1a | 2478 seconds |
|
||
|
|
| Stage 1b | 12935 seconds |
|
||
|
|
| Stage 2 | 36192 seconds |
|
||
|
|
|
||
|
|
Total recorded training runtime was approximately 14.3 hours on a local CUDA GPU.
|