k / step2_finetune.py
EmiliaLee's picture
Upload step2_finetune.py
b02cae0 verified
Raw
History Blame Contribute Delete
4.45 kB
"""Step 2: Fine-tune dense retriever with PQ training data.
Uses sentence-transformers MultipleNegativesRankingLoss with hard negatives.
Usage:
uv run python experiments/exp_021_enrichment_distillation/scripts/step2_finetune.py \
--domain wands --epochs 1 --batch_size 32
# Quick test run
uv run python experiments/exp_021_enrichment_distillation/scripts/step2_finetune.py \
--domain wands --epochs 1 --batch_size 32 --max_pairs 5000
"""
import argparse
import json
import sys
from pathlib import Path
import torch
from datasets import Dataset as HFDataset
from sentence_transformers import (
SentenceTransformer,
SentenceTransformerTrainer,
SentenceTransformerTrainingArguments,
)
from sentence_transformers.losses import MultipleNegativesRankingLoss
DATA_DIR = Path(__file__).parent.parent / "data"
MODEL_DIR = Path(__file__).parent.parent / "models"
def load_pairs(path: Path, max_pairs: int | None = None) -> list[dict]:
pairs = []
with open(path) as f:
for line in f:
if line.strip():
d = json.loads(line)
if d["negatives"]:
pairs.append(d)
if max_pairs and len(pairs) >= max_pairs:
break
return pairs
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--domain", required=True)
parser.add_argument("--model_name", default="BAAI/bge-m3")
parser.add_argument("--epochs", type=int, default=1)
parser.add_argument("--batch_size", type=int, default=32)
parser.add_argument("--lr", type=float, default=2e-5)
parser.add_argument("--warmup_ratio", type=float, default=0.1)
parser.add_argument("--max_pairs", type=int, default=None,
help="Limit training pairs (for quick test)")
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()
device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
print(f"Device: {device}")
# Load training data
train_path = DATA_DIR / f"{args.domain}_pq_train.jsonl"
print(f"Loading training data from {train_path}...")
pairs = load_pairs(train_path, args.max_pairs)
print(f" {len(pairs)} training pairs")
# Split: 95% train, 5% eval
split_idx = int(len(pairs) * 0.95)
train_pairs = pairs[:split_idx]
eval_pairs = pairs[split_idx:]
print(f" Train: {len(train_pairs)}, Eval: {len(eval_pairs)}")
train_dataset = HFDataset.from_dict({
"anchor": [p["query"] for p in train_pairs],
"positive": [p["positive"] for p in train_pairs],
"negative": [p["negatives"][0] if p["negatives"] else "" for p in train_pairs],
})
eval_dataset = HFDataset.from_dict({
"anchor": [p["query"] for p in eval_pairs],
"positive": [p["positive"] for p in eval_pairs],
"negative": [p["negatives"][0] if p["negatives"] else "" for p in eval_pairs],
})
# Load model
print(f"Loading {args.model_name}...")
model = SentenceTransformer(args.model_name, device=device)
# Loss
loss = MultipleNegativesRankingLoss(model)
# Output dir
output_dir = MODEL_DIR / f"{args.domain}_pq_finetuned"
output_dir.mkdir(parents=True, exist_ok=True)
# Training args
training_args = SentenceTransformerTrainingArguments(
output_dir=str(output_dir),
num_train_epochs=args.epochs,
per_device_train_batch_size=args.batch_size,
per_device_eval_batch_size=args.batch_size,
learning_rate=args.lr,
warmup_steps=args.warmup_ratio,
fp16=False,
bf16=(device == "cuda"),
eval_strategy="steps",
eval_steps=500,
save_strategy="epoch",
save_total_limit=2,
logging_steps=100,
seed=args.seed,
dataloader_num_workers=4 if device == "cuda" else 0,
report_to="none",
)
# Train
trainer = SentenceTransformerTrainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
loss=loss,
)
print(f"Training... (epochs={args.epochs}, batch_size={args.batch_size}, lr={args.lr})")
trainer.train()
# Save final model
final_path = MODEL_DIR / f"{args.domain}_pq_finetuned_final"
model.save(str(final_path))
print(f"\nModel saved → {final_path}")
if __name__ == "__main__":
main()