GoktugD's picture
Publish measured GöktuğTR release
a9e934f verified
Raw
History Blame Contribute Delete
8.04 kB
#!/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()