multilingual_model / train_sft.py
hjovi1's picture
Upload 6 files
f908c84 verified
Raw
History Blame Contribute Delete
4.64 kB
"""
SFT training for multilingual_model — Helena
Phase 1: fine-tune Qwen3-1.7B on multilingual MC data (belebele + xcopa)
to teach the model to output \boxed{letter} answers.
Same LoRA setup as trainingPhuc/train_sft.py — keeping it consistent
across the team so we can compare results.
Usage:
python -m trainingHelena.train_sft \
--data-dir /scratch/multilingual_data \
--output-dir /scratch/multilingual_model_sft
"""
import argparse
from pathlib import Path
import torch
from datasets import load_dataset
from peft import LoraConfig
from transformers import AutoModelForCausalLM, AutoTokenizer
from trl import SFTConfig, SFTTrainer
# LoRA target modules for Qwen3 (same as math model)
LORA_TARGETS = ["q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj"]
def get_args():
parser = argparse.ArgumentParser()
parser.add_argument("--base-model", default="/shared-ro/models/Qwen/Qwen3-1.7B")
parser.add_argument("--data-dir", default="/scratch/multilingual_data")
parser.add_argument("--output-dir", default="/scratch/multilingual_model_sft")
parser.add_argument("--epochs", type=int, default=3)
parser.add_argument("--batch-size", type=int, default=2)
parser.add_argument("--grad-accum", type=int, default=8)
parser.add_argument("--lr", type=float, default=2e-5)
parser.add_argument("--lora-r", type=int, default=64)
parser.add_argument("--max-seq-len", type=int, default=1024)
parser.add_argument("--report-to", default="none")
parser.add_argument("--run-name", default="multilingual-sft")
parser.add_argument("--resume", default=None,
help="path to checkpoint to resume from")
return parser.parse_args()
def main():
args = get_args()
output_dir = Path(args.output_dir)
merged_dir = output_dir / "merged"
# load tokenizer
tokenizer = AutoTokenizer.from_pretrained(args.base_model)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
# load model in bf16 (A100 supports it natively)
print(f"Loading base model from {args.base_model}...")
model = AutoModelForCausalLM.from_pretrained(
args.base_model,
torch_dtype=torch.bfloat16,
device_map="auto",
)
# load datasets
data_dir = Path(args.data_dir)
ds = load_dataset("json", data_files={
"train": str(data_dir / "sft_train.jsonl"),
"validation": str(data_dir / "sft_val.jsonl"),
})
print(f"Train: {len(ds['train'])} examples, Val: {len(ds['validation'])} examples")
# LoRA config — r=64 as in the math model
lora_config = LoraConfig(
r=args.lora_r,
lora_alpha=args.lora_r * 2, # standard: alpha = 2 * r
target_modules=LORA_TARGETS,
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
)
# SFT training
# 3 epochs because our dataset is much smaller than NuminaMath
trainer = SFTTrainer(
model=model,
processing_class=tokenizer,
train_dataset=ds["train"],
eval_dataset=ds["validation"],
peft_config=lora_config,
args=SFTConfig(
output_dir=str(output_dir),
num_train_epochs=args.epochs,
per_device_train_batch_size=args.batch_size,
gradient_accumulation_steps=args.grad_accum,
learning_rate=args.lr,
warmup_ratio=0.05,
lr_scheduler_type="cosine",
bf16=True,
max_length=args.max_seq_len,
dataset_text_field="text",
logging_steps=20,
eval_strategy="steps",
eval_steps=100,
save_strategy="steps",
save_steps=100,
save_total_limit=2,
load_best_model_at_end=True,
report_to=args.report_to,
run_name=args.run_name,
),
)
print("Starting SFT training...")
trainer.train(resume_from_checkpoint=args.resume)
# merge LoRA weights back into the base model so we get a standalone checkpoint
print(f"Merging LoRA weights → {merged_dir}")
merged_dir.mkdir(parents=True, exist_ok=True)
merged = trainer.model.merge_and_unload()
merged.save_pretrained(str(merged_dir))
tokenizer.save_pretrained(str(merged_dir))
print(f"Merged checkpoint saved: {merged_dir}")
print(f"\nNext: python -m trainingHelena.push_to_hub --checkpoint {merged_dir} --repo cs-552-2026-eminem-p/multilingual_model")
if __name__ == "__main__":
main()