"""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()