| import logging | |
| import sys | |
| from datetime import datetime | |
| import torch | |
| from datasets import load_dataset, DatasetDict | |
| from sentence_transformers import SentenceTransformer, losses | |
| from sentence_transformers.evaluation import ( | |
| TripletEvaluator, | |
| EmbeddingSimilarityEvaluator, | |
| SequentialEvaluator, | |
| ) | |
| from sentence_transformers.trainer import SentenceTransformerTrainer | |
| from sentence_transformers.training_args import ( | |
| SentenceTransformerTrainingArguments, | |
| BatchSamplers, | |
| MultiDatasetBatchSamplers, | |
| ) | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Logging | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| logging.basicConfig( | |
| format="%(asctime)s - %(message)s", | |
| datefmt="%Y-%m-%d %H:%M:%S", | |
| level=logging.INFO, | |
| handlers=[logging.FileHandler("logs.txt")], | |
| ) | |
| class Tee: | |
| def __init__(self, *files): self.files = files | |
| def write(self, obj): | |
| for f in self.files: f.write(obj); f.flush() | |
| def flush(self): | |
| for f in self.files: f.flush() | |
| def isatty(self): return False | |
| sys.stdout = Tee(sys.stdout, open("logs.txt", "a")) | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Config | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| train_batch_size = 32 | |
| model_name = "bert-base-arabertv02" | |
| model_nickname = "arabert" | |
| timestamp = datetime.now().strftime("%Y%m%d_%H%M") | |
| output_dir = f"output/{model_nickname}_{timestamp}" | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| matryoshka_dims = [768, 512, 256, 128, 64] | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # 1. Model | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| model = SentenceTransformer(model_name, device=device) | |
| model.set_pooling_include_prompt(include_prompt=False) | |
| print(f"Model running on: {device}") | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # TEST MODE β set to True to run on tiny slices and verify the | |
| # pipeline works end-to-end before full training | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| TEST_MODE = False | |
| TEST_SAMPLES = 100 | |
| def maybe_slice(dataset): | |
| if TEST_MODE: | |
| return dataset.select(range(min(TEST_SAMPLES, len(dataset)))) | |
| return dataset | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # 2. Datasets | |
| # | |
| # multineg_4_ss β SS | anchor, positive, neg_1..neg_4 | MNR | MultiNeg_4_ss.csv | |
| # multineg_30_ss β SS | anchor, positive, neg_1..neg_30 | MNR | MultiNeg_30_ss.csv | |
| # contrastive_ss β SS | sentence1, sentence2, label | Contrastive | s1_s2_label_ss.csv | |
| # contrastive_sts β STS | sentence1, sentence2, label | Contrastive | s1_s2_label_sts.csv | |
| # ap_ss β SS | anchor, positive | MNR | a_p_ss.csv | |
| # cosent_sts β STS | sentence1, sentence2, score | CoSENT | s1_s2_score_sts.csv | |
| # apn_ss β SS | anchor, positive, negative | MNR | a_p_n_ss.csv | |
| # apn_sts β STS | anchor, positive, negative | MNR | a_p_n_sts.csv | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| logging.info("Loading datasets...") | |
| def split_dataset(path, test_size=0.05, seed=42): | |
| """Load a CSV and split 95/5 into train and val in one shot.""" | |
| full = load_dataset("csv", data_files=path)["train"] | |
| if TEST_MODE: | |
| full = full.select(range(min(TEST_SAMPLES * 2, len(full)))) | |
| splits = full.train_test_split(test_size=test_size, seed=seed) | |
| return splits["train"], splits["test"] | |
| def load_train_only(path): | |
| """Load a CSV that has no evaluator β only needs a train split.""" | |
| full = load_dataset("csv", data_files=path)["train"] | |
| if TEST_MODE: | |
| full = full.select(range(min(TEST_SAMPLES, len(full)))) | |
| return full | |
| # Datasets only used for training (no evaluator needs them) | |
| multineg_4_train, _ = split_dataset("MultiNeg_4_ss.csv") | |
| multineg_30_train, _ = split_dataset("MultiNeg_30_ss.csv") | |
| contrastive_ss_train, _ = split_dataset("s1_s2_label_ss.csv") | |
| contrastive_sts_train, _ = split_dataset("s1_s2_label_sts.csv") | |
| ap_ss_train, _ = split_dataset("a_p_ss.csv") | |
| # Datasets split into train + val (evaluators use the val portion) | |
| apn_ss_train, apn_ss_val = split_dataset("a_p_n_ss.csv") | |
| cosent_sts_train, cosent_sts_val = split_dataset("s1_s2_score_sts.csv") | |
| apn_sts_train, apn_sts_val = split_dataset("a_p_n_sts.csv") | |
| train_dataset = DatasetDict({ | |
| "multineg_4_ss": multineg_4_train, | |
| "multineg_30_ss": multineg_30_train, | |
| "contrastive_ss": contrastive_ss_train, | |
| "contrastive_sts": contrastive_sts_train, | |
| "ap_ss": ap_ss_train, | |
| "cosent_sts": cosent_sts_train, | |
| "apn_ss": apn_ss_train, | |
| "apn_sts": apn_sts_train, | |
| }) | |
| # Val: only the 3 datasets that have evaluators | |
| eval_dataset = DatasetDict({ | |
| "apn_ss": apn_ss_val, | |
| "cosent_sts": cosent_sts_val, | |
| "apn_sts": apn_sts_val, | |
| }) | |
| logging.info(train_dataset) | |
| logging.info(eval_dataset) | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # 3. Loss functions β each wrapped in MatryoshkaLoss | |
| # Keys must exactly match the DatasetDict keys above | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def matryoshka(inner_loss): | |
| return losses.MatryoshkaLoss(model, inner_loss, matryoshka_dims=matryoshka_dims) | |
| loss = { | |
| "multineg_4_ss": matryoshka(losses.MultipleNegativesRankingLoss(model)), | |
| "multineg_30_ss": matryoshka(losses.MultipleNegativesRankingLoss(model)), | |
| "contrastive_ss": matryoshka(losses.ContrastiveLoss(model)), | |
| "contrastive_sts": matryoshka(losses.ContrastiveLoss(model)), | |
| "ap_ss": matryoshka(losses.MultipleNegativesRankingLoss(model)), | |
| "cosent_sts": matryoshka(losses.CoSENTLoss(model)), | |
| "apn_ss": matryoshka(losses.MultipleNegativesRankingLoss(model)), | |
| "apn_sts": matryoshka(losses.MultipleNegativesRankingLoss(model)), | |
| } | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # 4. Prompts | |
| # Format used: Dict[dataset_name, Dict[column_name, prompt]] | |
| # This is format 4 from the official docs β most granular and | |
| # correct for multi-task training where SS β STS prompts. | |
| # Matches exactly the pattern from training_nq_prompts.py where | |
| # prompts are passed via SentenceTransformerTrainingArguments. | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| ss_query_prompt = "Ω Ψ«ΩΩ ΩΨ°Ψ§ Ψ§ΩΨ³Ψ€Ψ§Ω Ψ§ΩΨΉΨ±Ψ¨Ω ΩΩΨ¨ΨΨ« ΨΉΩ Ψ§ΩΩ ΩΨ§Ψ·ΨΉ Ψ°Ψ§Ψͺ Ψ§ΩΨ΅ΩΨ©: " | |
| ss_passage_prompt = "Ω Ψ«ΩΩ ΩΨ°Ψ§ Ψ§ΩΩ ΩΨ·ΨΉ Ψ§ΩΨΉΨ±Ψ¨Ω ΩΩΨ§Ψ³ΨͺΨ±Ψ¬Ψ§ΨΉ: " | |
| sts_prompt = "Ω Ψ«ΩΩ ΩΨ°Ω Ψ§ΩΨ¬Ω ΩΨ© Ψ§ΩΨΉΨ±Ψ¨ΩΨ© ΩΩΨͺΨ΄Ψ§Ψ¨Ω Ψ§ΩΨ―ΩΨ§ΩΩ: " | |
| prompts = { | |
| # SS datasets β asymmetric: query side != passage side | |
| "multineg_4_ss": { | |
| "anchor": ss_query_prompt, | |
| "positive": ss_passage_prompt, | |
| "negative_1": ss_passage_prompt, | |
| "negative_2": ss_passage_prompt, | |
| "negative_3": ss_passage_prompt, | |
| "negative_4": ss_passage_prompt, | |
| }, | |
| "multineg_30_ss": { | |
| "anchor": ss_query_prompt, | |
| "positive": ss_passage_prompt, | |
| **{f"negative_{i}": ss_passage_prompt for i in range(1, 6)}, | |
| }, | |
| "contrastive_ss": { | |
| "sentence1": ss_query_prompt, | |
| "sentence2": ss_passage_prompt, | |
| }, | |
| "ap_ss": { | |
| "anchor": ss_query_prompt, | |
| "positive": ss_passage_prompt, | |
| }, | |
| "apn_ss": { | |
| "anchor": ss_query_prompt, | |
| "positive": ss_passage_prompt, | |
| "negative": ss_passage_prompt, | |
| }, | |
| # STS datasets β symmetric: both sides get the same prompt | |
| "contrastive_sts": { | |
| "sentence1": sts_prompt, | |
| "sentence2": sts_prompt, | |
| }, | |
| "cosent_sts": { | |
| "sentence1": sts_prompt, | |
| "sentence2": sts_prompt, | |
| }, | |
| "apn_sts": { | |
| "anchor": sts_prompt, | |
| "positive": sts_prompt, | |
| "negative": sts_prompt, | |
| }, | |
| } | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # 5. Evaluators | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def make_triplet_evaluators(dataset, name_prefix, max_samples=3_000): | |
| sample = dataset.shuffle(seed=42).select(range(min(max_samples, len(dataset)))) | |
| return [ | |
| TripletEvaluator( | |
| anchors=sample["anchor"], | |
| positives=sample["positive"], | |
| negatives=sample["negative"], | |
| name=f"{name_prefix}-{dim}", | |
| truncate_dim=dim, | |
| ) | |
| for dim in matryoshka_dims | |
| ] | |
| def make_sts_evaluators(dataset, name_prefix, max_samples=3_000): | |
| sample = dataset.shuffle(seed=42).select(range(min(max_samples, len(dataset)))) | |
| return [ | |
| EmbeddingSimilarityEvaluator( | |
| sentences1=sample["sentence1"], | |
| sentences2=sample["sentence2"], | |
| scores=sample["score"], | |
| name=f"{name_prefix}-{dim}", | |
| truncate_dim=dim, | |
| ) | |
| for dim in matryoshka_dims | |
| ] | |
| dev_evaluator = SequentialEvaluator( | |
| make_triplet_evaluators(eval_dataset["apn_ss"], "val-ss-apn") + | |
| make_sts_evaluators(eval_dataset["cosent_sts"], "val-sts-cosent") + | |
| make_triplet_evaluators(eval_dataset["apn_sts"], "val-sts-apn"), | |
| main_score_function=lambda scores: scores[0], # SS triplet at full dim is the main metric | |
| ) | |
| logging.info("Pre-training evaluation:") | |
| dev_evaluator(model) | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # 6. Training Arguments | |
| # prompts= is passed here, exactly as in the official example | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| args = SentenceTransformerTrainingArguments( | |
| output_dir=output_dir, | |
| seed=42, | |
| num_train_epochs=2, | |
| per_device_train_batch_size=train_batch_size, | |
| per_device_eval_batch_size=train_batch_size, | |
| gradient_accumulation_steps=2, | |
| bf16=True, | |
| fp16=False, | |
| learning_rate=2e-5, | |
| lr_scheduler_type="linear", | |
| warmup_ratio=0.1, | |
| weight_decay=0.01, | |
| batch_sampler=BatchSamplers.NO_DUPLICATES, | |
| multi_dataset_batch_sampler=MultiDatasetBatchSamplers.PROPORTIONAL, # sample proportionally to dataset size | |
| dataloader_num_workers=8, | |
| eval_strategy="steps", | |
| eval_steps=12000, | |
| save_strategy="steps", | |
| save_steps=12000, | |
| save_total_limit=2, | |
| report_to="tensorboard", | |
| logging_steps=200, | |
| logging_dir=f"{output_dir}/runs", | |
| prompts=prompts, # β official way to pass prompts, per dataset per column | |
| ) | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # 7. Trainer | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| trainer = SentenceTransformerTrainer( | |
| model=model, | |
| args=args, | |
| train_dataset=train_dataset, | |
| eval_dataset=eval_dataset, | |
| loss=loss, | |
| evaluator=dev_evaluator, | |
| ) | |
| trainer.train() | |
| print('finished') | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # 8. Save β store prompts in model config for clean inference | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| final_output_dir = f"{output_dir}/final" | |
| # Save the prompts in the model config so users can call model.encode(..., prompt_name="query") at inference | |
| model.prompts = { | |
| "ss_query": ss_query_prompt, | |
| "ss_passage": ss_passage_prompt, | |
| "sts": sts_prompt, | |
| } | |
| model.save(final_output_dir) | |
| print(f"Model saved to {final_output_dir}") | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # 9. Benchmark test evaluation β run ONCE, never during training | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| #test_triplet = load_dataset("csv", data_files="benchmark_test_triplet.csv")["train"] | |
| #test_sts = load_dataset("csv", data_files="benchmark_test_sts.csv")["train"] | |
| #test_evaluator = SequentialEvaluator( | |
| # make_triplet_evaluators(test_triplet, "test-benchmark-ss") + | |
| # make_sts_evaluators(test_sts, "test-benchmark-sts") | |
| #) | |
| #results = test_evaluator(model, output_path=final_output_dir) | |
| #print("Benchmark test results:", results) | |