rndubs's picture
Add scripts/sft_train.py
dbcbd01 verified
Raw
History Blame Contribute Delete
9.58 kB
#!/usr/bin/env python3
"""
============================================================================
Supervised Fine-Tuning (SFT) with TRL
============================================================================
After pre-training with Megatron-LM and converting to HuggingFace format,
this script performs SFT using TRL's SFTTrainer.
SFT teaches the model to follow instructions, write code on command,
use Slurm, and reason step-by-step.
Usage:
# Single-node 4×H100
torchrun --nproc_per_node=4 scripts/sft_train.py \
--model-path /path/to/hf-model \
--output-dir /path/to/sft-model \
--dataset-path /path/to/sft-data.jsonl
# Multi-node via Slurm (see slurm/sft.sbatch)
Prerequisites:
pip install transformers trl datasets accelerate peft torch
============================================================================
"""
import argparse
import json
import os
import sys
from pathlib import Path
import torch
from datasets import load_dataset, Dataset, concatenate_datasets
def prepare_sft_dataset(data_paths: list[str], max_samples: int = None) -> Dataset:
"""
Load and prepare SFT datasets.
Expects data in ChatML/messages format:
{"messages": [{"role": "system", "content": "..."}, {"role": "user", "content": "..."}, {"role": "assistant", "content": "..."}]}
"""
all_datasets = []
for path in data_paths:
if path.startswith("hf://") or "/" in path and not os.path.exists(path):
# HuggingFace Hub dataset
hub_name = path.replace("hf://", "")
print(f"Loading HF dataset: {hub_name}")
ds = load_dataset(hub_name, split="train")
else:
# Local JSONL file
print(f"Loading local dataset: {path}")
ds = load_dataset("json", data_files=path, split="train")
# Ensure messages format
if "messages" in ds.column_names:
all_datasets.append(ds)
elif "instruction" in ds.column_names and "output" in ds.column_names:
# Convert instruction/output format to messages
def convert_to_messages(example):
messages = []
if example.get("system"):
messages.append({"role": "system", "content": example["system"]})
messages.append({"role": "user", "content": example["instruction"]})
if example.get("input"):
messages[-1]["content"] += f"\n\n{example['input']}"
messages.append({"role": "assistant", "content": example["output"]})
return {"messages": messages}
ds = ds.map(convert_to_messages, remove_columns=ds.column_names)
all_datasets.append(ds)
elif "prompt" in ds.column_names and "completion" in ds.column_names:
# Convert prompt/completion format
def convert_prompt_completion(example):
messages = [
{"role": "user", "content": example["prompt"]},
{"role": "assistant", "content": example["completion"]},
]
return {"messages": messages}
ds = ds.map(convert_prompt_completion, remove_columns=ds.column_names)
all_datasets.append(ds)
else:
print(f" WARNING: Unknown format in {path}, columns: {ds.column_names}")
continue
print(f" Loaded {len(ds)} examples")
if not all_datasets:
print("ERROR: No valid datasets loaded!")
sys.exit(1)
combined = concatenate_datasets(all_datasets)
if max_samples:
combined = combined.shuffle(seed=42).select(range(min(max_samples, len(combined))))
print(f"\nTotal SFT examples: {len(combined)}")
return combined
def main():
parser = argparse.ArgumentParser(description="SFT Training with TRL")
parser.add_argument("--model-path", required=True, help="Path to pre-trained HF model")
parser.add_argument("--output-dir", required=True, help="Output directory for SFT model")
parser.add_argument("--dataset-path", nargs="+", required=True, help="SFT dataset paths (local JSONL or HF Hub)")
parser.add_argument("--max-seq-length", type=int, default=8192)
parser.add_argument("--num-epochs", type=int, default=3)
parser.add_argument("--learning-rate", type=float, default=2e-5)
parser.add_argument("--per-device-batch-size", type=int, default=2)
parser.add_argument("--gradient-accumulation-steps", type=int, default=8)
parser.add_argument("--warmup-ratio", type=float, default=0.03)
parser.add_argument("--max-samples", type=int, default=None)
parser.add_argument("--push-to-hub", action="store_true")
parser.add_argument("--hub-model-id", type=str, default=None)
parser.add_argument("--use-lora", action="store_true", help="Use LoRA for parameter-efficient SFT")
parser.add_argument("--lora-rank", type=int, default=64)
args = parser.parse_args()
# ================================================================
# Load tokenizer and model
# ================================================================
from transformers import AutoTokenizer, AutoModelForCausalLM
print(f"Loading model: {args.model_path}")
tokenizer = AutoTokenizer.from_pretrained(args.model_path, trust_remote_code=True)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
model_kwargs = {
"trust_remote_code": True,
"torch_dtype": torch.bfloat16,
"attn_implementation": "flash_attention_2", # requires flash-attn
}
if args.use_lora:
# Load in 4-bit for LoRA
from transformers import BitsAndBytesConfig
model_kwargs["quantization_config"] = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
)
model = AutoModelForCausalLM.from_pretrained(args.model_path, **model_kwargs)
# ================================================================
# LoRA config (optional)
# ================================================================
peft_config = None
if args.use_lora:
from peft import LoraConfig
peft_config = LoraConfig(
r=args.lora_rank,
lora_alpha=args.lora_rank * 2,
lora_dropout=0.05,
target_modules=[
"q_proj", "k_proj", "v_proj", "o_proj", # attention
"gate_proj", "up_proj", "down_proj", # experts
],
task_type="CAUSAL_LM",
)
print(f"Using LoRA with rank={args.lora_rank}")
# ================================================================
# Load dataset
# ================================================================
dataset = prepare_sft_dataset(args.dataset_path, args.max_samples)
# ================================================================
# Training config
# ================================================================
from trl import SFTConfig, SFTTrainer
training_args = SFTConfig(
output_dir=args.output_dir,
num_train_epochs=args.num_epochs,
per_device_train_batch_size=args.per_device_batch_size,
gradient_accumulation_steps=args.gradient_accumulation_steps,
learning_rate=args.learning_rate,
lr_scheduler_type="cosine",
warmup_ratio=args.warmup_ratio,
weight_decay=0.01,
max_grad_norm=1.0,
bf16=True,
max_seq_length=args.max_seq_length,
packing=True, # Pack multiple examples into one sequence for efficiency
logging_strategy="steps",
logging_steps=10,
logging_first_step=True,
disable_tqdm=True,
save_strategy="steps",
save_steps=500,
save_total_limit=3,
eval_strategy="steps",
eval_steps=500,
gradient_checkpointing=True,
gradient_checkpointing_kwargs={"use_reentrant": False},
report_to=["tensorboard"], # or ["wandb"]
push_to_hub=args.push_to_hub,
hub_model_id=args.hub_model_id,
seed=42,
dataloader_num_workers=4,
dataloader_pin_memory=True,
# DeepSpeed config for multi-GPU
deepspeed=None, # Use accelerate config instead for multi-GPU
)
# ================================================================
# Train
# ================================================================
trainer = SFTTrainer(
model=model,
args=training_args,
train_dataset=dataset,
processing_class=tokenizer,
peft_config=peft_config,
)
print(f"\n{'='*60}")
print(f"Starting SFT Training")
print(f"{'='*60}")
print(f"Model: {args.model_path}")
print(f"Dataset: {len(dataset)} examples")
print(f"Epochs: {args.num_epochs}")
print(f"LR: {args.learning_rate}")
print(f"Batch: {args.per_device_batch_size} × {args.gradient_accumulation_steps} × {torch.cuda.device_count()} GPUs")
print(f"Max seq length: {args.max_seq_length}")
print(f"LoRA: {args.use_lora} (rank={args.lora_rank if args.use_lora else 'N/A'})")
print(f"Output: {args.output_dir}")
print(f"{'='*60}\n")
trainer.train()
# Save
trainer.save_model()
tokenizer.save_pretrained(args.output_dir)
if args.push_to_hub and args.hub_model_id:
trainer.push_to_hub()
print(f"\nSFT training complete! Model saved to: {args.output_dir}")
if __name__ == "__main__":
main()