#!/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()