#!/usr/bin/env python3 """Fine-tune a Turkish dense retriever with reproducible settings.""" from __future__ import annotations import argparse import json import platform import random import re import time from pathlib import Path import datasets import numpy as np import sentence_transformers import torch import transformers import yaml from datasets import load_dataset from sentence_transformers import ( SentenceTransformer, SentenceTransformerTrainer, SentenceTransformerTrainingArguments, losses, ) from sentence_transformers.evaluation import TripletEvaluator from sentence_transformers.training_args import BatchSamplers from goktugtr.text import training_query def seed_everything(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--config", type=Path, default=Path("configs/train_270m.yaml")) parser.add_argument("--data-dir", type=Path, default=Path("data/processed")) parser.add_argument("--output-dir", type=Path, default=Path("outputs/goktugtr-270m")) parser.add_argument("--max-train-rows", type=int) parser.add_argument("--max-steps", type=int, default=-1) return parser.parse_args() def latest_complete_checkpoint(output_dir: Path) -> Path | None: """Return the newest checkpoint that contains all trainer resume state.""" candidates: list[tuple[int, Path]] = [] for path in output_dir.glob("checkpoint-*"): match = re.fullmatch(r"checkpoint-(\d+)", path.name) if not match: continue required = ( "model.safetensors", "optimizer.pt", "scheduler.pt", "trainer_state.json", "rng_state.pth", ) if all((path / filename).is_file() for filename in required): candidates.append((int(match.group(1)), path)) return max(candidates, default=(0, None), key=lambda item: item[0])[1] def main() -> None: args = parse_args() config = yaml.safe_load(args.config.read_text(encoding="utf-8")) seed_everything(int(config["seed"])) args.output_dir.mkdir(parents=True, exist_ok=True) dataset = load_dataset( "parquet", data_files={ "train": str(args.data_dir / "train.parquet"), "validation": str(args.data_dir / "validation.parquet"), }, ) train = dataset["train"] if args.max_train_rows: train = train.select(range(min(args.max_train_rows, len(train)))) def add_prompt(row: dict[str, str]) -> dict[str, str]: return { "anchor": training_query(row["query"], config.get("query_style", "harrier")), "positive": row["positive"], "negative": row["negative"], } remove_columns = dataset["train"].column_names train = train.map(add_prompt, remove_columns=remove_columns, desc="Formatting train queries") validation = dataset["validation"].map( add_prompt, remove_columns=remove_columns, desc="Formatting validation queries" ) train = train.select_columns(["anchor", "positive", "negative"]) validation = validation.select_columns(["anchor", "positive", "negative"]) model_kwargs = {"dtype": torch.bfloat16} if bool(config["bf16"]) else {} processor_kwargs = {"padding_side": config.get("padding_side", "right")} model = SentenceTransformer( config["base_model"], revision=config.get("base_model_revision"), model_kwargs=model_kwargs, processor_kwargs=processor_kwargs, ) model.max_seq_length = int(config["max_seq_length"]) evaluator = TripletEvaluator( anchors=validation["anchor"], positives=validation["positive"], negatives=validation["negative"], name="goktugtr-validation", batch_size=8, show_progress_bar=True, ) baseline = evaluator(model, output_path=str(args.output_dir), epoch=0, steps=0) (args.output_dir / "baseline_triplet.json").write_text( json.dumps(baseline, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" ) training_args = SentenceTransformerTrainingArguments( output_dir=str(args.output_dir), num_train_epochs=float(config["epochs"]), max_steps=args.max_steps, per_device_train_batch_size=int(config["per_device_batch_size"]), per_device_eval_batch_size=8, gradient_accumulation_steps=int(config["gradient_accumulation_steps"]), learning_rate=float(config["learning_rate"]), warmup_steps=float(config["warmup_ratio"]), bf16=bool(config["bf16"]), tf32=True, gradient_checkpointing=bool(config["gradient_checkpointing"]), gradient_checkpointing_kwargs={"use_reentrant": False}, optim="adamw_torch_fused", batch_sampler=BatchSamplers.NO_DUPLICATES, eval_strategy="steps", eval_steps=int(config["eval_steps"]), save_strategy="steps", save_steps=int(config["save_steps"]), save_total_limit=2, logging_steps=int(config["logging_steps"]), dataloader_num_workers=2, dataloader_pin_memory=True, report_to="none", run_name=config["project_name"], seed=int(config["seed"]), ) if config.get("loss") == "cached_multiple_negatives_ranking": loss = losses.CachedMultipleNegativesRankingLoss( model, mini_batch_size=int(config.get("loss_mini_batch_size", 2)), scale=20.0, ) else: loss = losses.MultipleNegativesRankingLoss(model, scale=20.0) trainer = SentenceTransformerTrainer( model=model, args=training_args, train_dataset=train, eval_dataset=validation, loss=loss, evaluator=evaluator, ) resume_checkpoint = latest_complete_checkpoint(args.output_dir) if resume_checkpoint: print(f"Resuming from complete checkpoint: {resume_checkpoint}", flush=True) torch.cuda.reset_peak_memory_stats() training_started = time.perf_counter() train_output = trainer.train( resume_from_checkpoint=str(resume_checkpoint) if resume_checkpoint else None ) training_seconds = time.perf_counter() - training_started trainer.state.save_to_json(str(args.output_dir / "trainer_state.json")) final_dir = args.output_dir / "final" model.save_pretrained(str(final_dir), safe_serialization=True) final_metrics = evaluator(model, output_path=str(args.output_dir), epoch=1, steps=-1) (args.output_dir / "final_triplet.json").write_text( json.dumps(final_metrics, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" ) environment = { "python": platform.python_version(), "torch": torch.__version__, "transformers": transformers.__version__, "sentence_transformers": sentence_transformers.__version__, "datasets": datasets.__version__, "cuda": torch.version.cuda, "gpu": torch.cuda.get_device_name(0) if torch.cuda.is_available() else None, "gpu_total_memory_gb": ( round(torch.cuda.get_device_properties(0).total_memory / 2**30, 3) if torch.cuda.is_available() else None ), "config": config, "train_rows": len(train), "validation_rows": len(validation), "training_seconds": training_seconds, "resumed_from_checkpoint": ( str(resume_checkpoint) if resume_checkpoint is not None else None ), "training_metrics": train_output.metrics, "max_gpu_memory_gb": round(torch.cuda.max_memory_allocated() / 2**30, 3), } (args.output_dir / "environment.json").write_text( json.dumps(environment, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" ) print(json.dumps({"baseline": baseline, "final": final_metrics}, indent=2)) if __name__ == "__main__": main()