| |
| """ |
| ============================================================================ |
| 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): |
| |
| hub_name = path.replace("hf://", "") |
| print(f"Loading HF dataset: {hub_name}") |
| ds = load_dataset(hub_name, split="train") |
| else: |
| |
| print(f"Loading local dataset: {path}") |
| ds = load_dataset("json", data_files=path, split="train") |
|
|
| |
| if "messages" in ds.column_names: |
| all_datasets.append(ds) |
| elif "instruction" in ds.column_names and "output" in ds.column_names: |
| |
| 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: |
| |
| 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() |
|
|
| |
| |
| |
| 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", |
| } |
|
|
| if args.use_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) |
|
|
| |
| |
| |
| 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", |
| "gate_proj", "up_proj", "down_proj", |
| ], |
| task_type="CAUSAL_LM", |
| ) |
| print(f"Using LoRA with rank={args.lora_rank}") |
|
|
| |
| |
| |
| dataset = prepare_sft_dataset(args.dataset_path, args.max_samples) |
|
|
| |
| |
| |
| 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, |
| 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"], |
| 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=None, |
| ) |
|
|
| |
| |
| |
| 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() |
|
|
| |
| 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() |
|
|