| """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}") |
|
|
| |
| 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_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], |
| }) |
|
|
| |
| print(f"Loading {args.model_name}...") |
| model = SentenceTransformer(args.model_name, device=device) |
|
|
| |
| loss = MultipleNegativesRankingLoss(model) |
|
|
| |
| output_dir = MODEL_DIR / f"{args.domain}_pq_finetuned" |
| output_dir.mkdir(parents=True, exist_ok=True) |
|
|
| |
| 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", |
| ) |
|
|
| |
| 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() |
|
|
| |
| 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() |
|
|