Qwen-text2SQL / train.py
AaronTekle's picture
Update train.py
cb2dfdf verified
Raw
History Blame Contribute Delete
16.7 kB
import inspect
import math
import os
import torch
from datasets import load_dataset
from peft import LoraConfig
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
set_seed,
)
from trl import SFTConfig, SFTTrainer
from config import (
DATASET_ID,
MAX_TRAIN_SAMPLES,
MAX_VALIDATION_SAMPLES,
MODEL_ID,
OUTPUT_DIR,
SEED,
SYSTEM_PROMPT,
)
# ============================================================
# TRAINING CONFIGURATION
# ============================================================
NUM_EPOCHS = 2
TRAIN_BATCH_SIZE = 4
EVAL_BATCH_SIZE = 4
GRADIENT_ACCUMULATION_STEPS = 4
LEARNING_RATE = 2e-4
MAX_SEQUENCE_LENGTH = 2048
WARMUP_PERCENT = 0.03
# ============================================================
# DATA FORMATTING
# ============================================================
def format_for_sft(example):
"""
Convert one Gretel Text-to-SQL record into
conversational prompt-completion format.
Input:
Database schema/context
+
Natural-language request
Target:
SQL query
"""
user_message = (
"Database context:\n"
f"{example['sql_context']}\n\n"
"Request:\n"
f"{example['sql_prompt']}"
)
return {
"prompt": [
{
"role": "system",
"content": SYSTEM_PROMPT,
},
{
"role": "user",
"content": user_message,
},
],
"completion": [
{
"role": "assistant",
"content": example["sql"].strip(),
},
],
}
# ============================================================
# CUDA / RTX CHECK
# ============================================================
def configure_gpu():
"""
Require an NVIDIA CUDA GPU.
This intentionally DOES NOT fall back to CPU.
Training a 0.5B model on CPU would be extremely slow.
"""
print("=" * 65)
print("GPU CHECK")
print("=" * 65)
print("PyTorch version:", torch.__version__)
print("PyTorch CUDA runtime:", torch.version.cuda)
print("CUDA available:", torch.cuda.is_available())
if not torch.cuda.is_available():
raise RuntimeError(
"\nCUDA is NOT available.\n\n"
"Your train.py is configured for an NVIDIA GeForce RTX GPU, "
"but your current PyTorch installation cannot access CUDA.\n\n"
"If your PyTorch version contains '+cpu', you installed the "
"CPU-only version of PyTorch.\n\n"
"Run `nvidia-smi` first to verify Windows sees your GPU, "
"then install the CUDA-enabled version of PyTorch."
)
# NVIDIA GPU
device = torch.device("cuda:0")
torch.cuda.set_device(device)
gpu_name = torch.cuda.get_device_name(device)
properties = torch.cuda.get_device_properties(device)
total_vram_gb = (
properties.total_memory
/ (1024 ** 3)
)
print("GPU:", gpu_name)
print(
"VRAM:",
f"{total_vram_gb:.2f} GB"
)
print(
"CUDA device:",
torch.cuda.current_device()
)
print("=" * 65)
return device
# ============================================================
# SELECT TRAINING PRECISION
# ============================================================
def select_precision():
"""
Automatically select BF16 or FP16.
BF16:
Preferred when the GPU supports it.
FP16:
Used otherwise on CUDA GPUs.
"""
bf16_supported = (
torch.cuda.is_bf16_supported()
)
if bf16_supported:
model_dtype = torch.bfloat16
use_bf16 = True
use_fp16 = False
precision_name = "BF16"
else:
model_dtype = torch.float16
use_bf16 = False
use_fp16 = True
precision_name = "FP16"
print(
"Training precision:",
precision_name
)
print(
"Model dtype:",
model_dtype
)
return (
model_dtype,
use_bf16,
use_fp16,
)
# ============================================================
# BUILD SFT CONFIG SAFELY
# ============================================================
def build_sft_config(
train_dataset,
use_bf16,
use_fp16,
):
"""
Build SFTConfig while checking which arguments
are supported by the installed TRL version.
This prevents errors such as:
unexpected keyword argument 'warmup_ratio'
caused by different TRL versions exposing slightly
different SFTConfig parameters.
"""
# --------------------------------------------------------
# Estimate optimizer steps
# --------------------------------------------------------
effective_batch_size = (
TRAIN_BATCH_SIZE
* GRADIENT_ACCUMULATION_STEPS
)
steps_per_epoch = math.ceil(
len(train_dataset)
/ effective_batch_size
)
total_training_steps = (
steps_per_epoch
* NUM_EPOCHS
)
warmup_steps = max(
1,
int(
total_training_steps
* WARMUP_PERCENT
),
)
print()
print("=" * 65)
print("TRAINING PLAN")
print("=" * 65)
print(
"Training samples:",
len(train_dataset)
)
print(
"Batch size:",
TRAIN_BATCH_SIZE
)
print(
"Gradient accumulation:",
GRADIENT_ACCUMULATION_STEPS
)
print(
"Effective batch size:",
effective_batch_size
)
print(
"Epochs:",
NUM_EPOCHS
)
print(
"Estimated optimizer steps:",
total_training_steps
)
print(
"Warmup steps:",
warmup_steps
)
print("=" * 65)
# --------------------------------------------------------
# Configuration
# --------------------------------------------------------
config_kwargs = {
"output_dir":
OUTPUT_DIR,
"num_train_epochs":
NUM_EPOCHS,
"per_device_train_batch_size":
TRAIN_BATCH_SIZE,
"per_device_eval_batch_size":
EVAL_BATCH_SIZE,
"gradient_accumulation_steps":
GRADIENT_ACCUMULATION_STEPS,
"learning_rate":
LEARNING_RATE,
"warmup_steps":
warmup_steps,
"weight_decay":
0.01,
"lr_scheduler_type":
"cosine",
"logging_steps":
10,
"eval_strategy":
"steps",
"eval_steps":
100,
"save_strategy":
"steps",
"save_steps":
100,
"save_total_limit":
2,
"load_best_model_at_end":
True,
"metric_for_best_model":
"eval_loss",
"greater_is_better":
False,
"bf16":
use_bf16,
"fp16":
use_fp16,
"max_length":
MAX_SEQUENCE_LENGTH,
# Prompt-completion datasets can calculate
# loss only on the SQL completion.
"completion_only_loss":
True,
"gradient_checkpointing":
True,
"report_to":
"none",
"seed":
SEED,
}
# --------------------------------------------------------
# TRL compatibility check
# --------------------------------------------------------
signature = inspect.signature(
SFTConfig.__init__
)
supported_parameters = (
signature.parameters
)
filtered_kwargs = {}
skipped_kwargs = []
for key, value in config_kwargs.items():
if key in supported_parameters:
filtered_kwargs[key] = value
else:
skipped_kwargs.append(key)
if skipped_kwargs:
print()
print(
"TRL compatibility notice:"
)
print(
"Unsupported SFTConfig arguments:",
skipped_kwargs,
)
print(
"They were skipped automatically."
)
return SFTConfig(
**filtered_kwargs
)
# ============================================================
# MAIN TRAINING FUNCTION
# ============================================================
def main():
# --------------------------------------------------------
# Reproducibility
# --------------------------------------------------------
set_seed(SEED)
# --------------------------------------------------------
# Require NVIDIA RTX GPU
# --------------------------------------------------------
device = configure_gpu()
# --------------------------------------------------------
# Select BF16 / FP16
# --------------------------------------------------------
(
model_dtype,
use_bf16,
use_fp16,
) = select_precision()
# --------------------------------------------------------
# Load Hugging Face dataset
# --------------------------------------------------------
print()
print("=" * 65)
print("LOADING DATASET")
print("=" * 65)
print(
"Dataset:",
DATASET_ID
)
dataset = load_dataset(
DATASET_ID
)
# Shuffle the original train/test splits.
train_dataset = (
dataset["train"]
.shuffle(seed=SEED)
)
validation_dataset = (
dataset["test"]
.shuffle(seed=SEED)
)
# --------------------------------------------------------
# Limit dataset during development
# --------------------------------------------------------
if MAX_TRAIN_SAMPLES is not None:
train_dataset = (
train_dataset.select(
range(
min(
MAX_TRAIN_SAMPLES,
len(train_dataset),
)
)
)
)
if MAX_VALIDATION_SAMPLES is not None:
validation_dataset = (
validation_dataset.select(
range(
min(
MAX_VALIDATION_SAMPLES,
len(validation_dataset),
)
)
)
)
print(
"Training examples:",
len(train_dataset)
)
print(
"Validation examples:",
len(validation_dataset)
)
# --------------------------------------------------------
# Format dataset for instruction tuning
# --------------------------------------------------------
print()
print("Formatting training data...")
train_dataset = train_dataset.map(
format_for_sft,
remove_columns=(
train_dataset.column_names
),
desc="Formatting training data",
)
print()
print("Formatting validation data...")
validation_dataset = (
validation_dataset.map(
format_for_sft,
remove_columns=(
validation_dataset.column_names
),
desc="Formatting validation data",
)
)
# --------------------------------------------------------
# Tokenizer
# --------------------------------------------------------
print()
print("=" * 65)
print("LOADING TOKENIZER")
print("=" * 65)
tokenizer = (
AutoTokenizer.from_pretrained(
MODEL_ID
)
)
if tokenizer.pad_token is None:
tokenizer.pad_token = (
tokenizer.eos_token
)
print(
"EOS token:",
tokenizer.eos_token
)
print(
"PAD token:",
tokenizer.pad_token
)
# --------------------------------------------------------
# Base model
# --------------------------------------------------------
print()
print("=" * 65)
print("LOADING QWEN BASE MODEL")
print("=" * 65)
print(
"Model:",
MODEL_ID
)
print(
"Target GPU:",
torch.cuda.get_device_name(0)
)
model = (
AutoModelForCausalLM
.from_pretrained(
MODEL_ID,
# Modern Transformers syntax.
dtype=model_dtype,
)
)
# Trainer / Accelerate handles moving the model
# onto the CUDA device.
# single-GPU training job. device_map is primarily
# useful for model dispatch/offloading.
model.config.use_cache = False
# --------------------------------------------------------
# Enable gradient checkpointing
# --------------------------------------------------------
model.gradient_checkpointing_enable()
# --------------------------------------------------------
# LoRA configuration
# --------------------------------------------------------
print()
print("=" * 65)
print("CONFIGURING LoRA")
print("=" * 65)
peft_config = LoraConfig(
# Rank of low-rank matrices A and B.
r=16,
# LoRA scaling parameter.
lora_alpha=32,
# Regularization on adapter branch.
lora_dropout=0.05,
# Do not train bias parameters.
bias="none",
# Qwen is a causal LM.
task_type="CAUSAL_LM",
# Apply LoRA to all four major
# self-attention projections.
target_modules=[
"q_proj",
"k_proj",
"v_proj",
"o_proj",
],
)
print("LoRA rank: 16")
print("LoRA alpha: 32")
print("LoRA dropout: 0.05")
print(
"LoRA targets:",
"q_proj, k_proj, v_proj, o_proj",
)
# --------------------------------------------------------
# Training configuration
# --------------------------------------------------------
training_args = build_sft_config(
train_dataset,
use_bf16,
use_fp16,
)
# --------------------------------------------------------
# Trainer
# --------------------------------------------------------
print()
print("=" * 65)
print("CREATING SFT TRAINER")
print("=" * 65)
trainer = SFTTrainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=validation_dataset,
processing_class=tokenizer,
peft_config=peft_config,
)
# --------------------------------------------------------
# Print parameter counts
# --------------------------------------------------------
trainable_parameters = 0
total_parameters = 0
for parameter in trainer.model.parameters():
parameter_count = (
parameter.numel()
)
total_parameters += (
parameter_count
)
if parameter.requires_grad:
trainable_parameters += (
parameter_count
)
trainable_percent = (
100
* trainable_parameters
/ total_parameters
)
print()
print("=" * 65)
print("PARAMETER SUMMARY")
print("=" * 65)
print(
f"Total parameters: "
f"{total_parameters:,}"
)
print(
f"Trainable parameters: "
f"{trainable_parameters:,}"
)
print(
f"Trainable percentage: "
f"{trainable_percent:.4f}%"
)
print()
print(
"Active GPU:",
torch.cuda.get_device_name(0)
)
print(
"Allocated VRAM:",
f"{torch.cuda.memory_allocated(0) / 1024**3:.2f} GB"
)
# --------------------------------------------------------
# Train
# --------------------------------------------------------
print()
print("=" * 65)
print("STARTING LoRA FINE-TUNING")
print("=" * 65)
trainer.train()
# --------------------------------------------------------
# Save trained LoRA adapter
# --------------------------------------------------------
print()
print("=" * 65)
print("SAVING MODEL")
print("=" * 65)
trainer.save_model(
OUTPUT_DIR
)
tokenizer.save_pretrained(
OUTPUT_DIR
)
print()
print(
"LoRA adapter saved to:",
OUTPUT_DIR
)
print(
"Final GPU:",
torch.cuda.get_device_name(0)
)
print(
"Maximum VRAM allocated:",
f"{torch.cuda.max_memory_allocated(0) / 1024**3:.2f} GB"
)
# ============================================================
# ENTRY POINT
# ============================================================
if __name__ == "__main__":
main()