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, 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}") checkpoint_path = "/home/skiredj.abderrahman/khalil/sbert_training/third_training/output/arabert_20260320_0436/checkpoint-84000/" # ───────────────────────────────────────────────────────────── # TEST MODE # ───────────────────────────────────────────────────────────── TEST_MODE = False TEST_SAMPLES = 100 # ───────────────────────────────────────────────────────────── # 2. Datasets # # multineg_4_ss → SS | anchor, positive, neg_1..neg_4 | MNR # multineg_30_ss → SS | anchor, positive, neg_1..neg_30 | MNR # contrastive_ss → SS | sentence1, sentence2, label | Contrastive # contrastive_sts → STS | sentence1, sentence2, label | Contrastive # ap_ss → SS | anchor, positive | MNR # cosent_sts → STS | sentence1, sentence2, score | CoSENT # apn_ss → SS | anchor, positive, negative | MNR # apn_sts → STS | anchor, positive, negative | MNR # ms_marco → SS | anchor, positive, negative | MNR (pre-split files) # ───────────────────────────────────────────────────────────── logging.info("Loading datasets...") def load_csv(path): ds = load_dataset("csv", data_files=path)["train"] # Drop rows with None, NaN, or empty string values ds = ds.filter(lambda x: all(x[col] is not None and str(x[col]).strip() != "" for col in ds.column_names)) if TEST_MODE: ds = ds.select(range(min(TEST_SAMPLES, len(ds)))) return ds multineg_4_train = load_csv("/home/skiredj.abderrahman/khalil/sbert_training/third_training/clean_data/MultiNeg_4_ss.csv") multineg_30_train = load_csv("/home/skiredj.abderrahman/khalil/sbert_training/third_training/clean_data/MultiNeg_30_ss.csv") contrastive_ss_train = load_csv("/home/skiredj.abderrahman/khalil/sbert_training/third_training/clean_data/s1_s2_label_ss.csv") contrastive_sts_train = load_csv("/home/skiredj.abderrahman/khalil/sbert_training/third_training/clean_data/s1_s2_label_sts.csv") ap_ss_train = load_csv("/home/skiredj.abderrahman/khalil/sbert_training/third_training/clean_data/a_p_ss.csv") cosent_sts_train = load_csv("/home/skiredj.abderrahman/khalil/sbert_training/third_training/clean_data/s1_s2_score_sts.csv") apn_ss_train = load_csv("/home/skiredj.abderrahman/khalil/sbert_training/third_training/clean_data/a_p_n_ss.csv") apn_sts_train = load_csv("/home/skiredj.abderrahman/khalil/sbert_training/third_training/clean_data/a_p_n_sts.csv") ms_marco_train = load_csv("/home/skiredj.abderrahman/khalil/sbert_training/third_training/clean_data/ms_marco_clean_dataset36_train.csv") ms_marco_val = load_csv("/home/skiredj.abderrahman/khalil/sbert_training/third_training/clean_data/ms_marco_clean_dataset36_val.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, "ms_marco": ms_marco_train, }) eval_dataset = DatasetDict({ "ms_marco": ms_marco_val, }) logging.info(train_dataset) logging.info(eval_dataset) # ───────────────────────────────────────────────────────────── # 3. Loss functions — each wrapped in MatryoshkaLoss # ───────────────────────────────────────────────────────────── 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)), "ms_marco": matryoshka(losses.MultipleNegativesRankingLoss(model)), } # ───────────────────────────────────────────────────────────── # 4. Evaluator — ms_marco_val only # ───────────────────────────────────────────────────────────── 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 ] dev_evaluator = SequentialEvaluator( make_triplet_evaluators(ms_marco_val, "val-ms-marco"), main_score_function=lambda scores: scores[0], ) logging.info("Pre-training evaluation:") dev_evaluator(model) # ───────────────────────────────────────────────────────────── # 5. Training Arguments # ───────────────────────────────────────────────────────────── 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, 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", ) # ───────────────────────────────────────────────────────────── # 6. Trainer # ───────────────────────────────────────────────────────────── trainer = SentenceTransformerTrainer( model=model, args=args, train_dataset=train_dataset, eval_dataset=eval_dataset, loss=loss, evaluator=dev_evaluator, ) print(f"continue from checkpoint path : {checkpoint_path}") trainer.train(resume_from_checkpoint=checkpoint_path) print("finished") # ───────────────────────────────────────────────────────────── # 7. Save # ───────────────────────────────────────────────────────────── final_output_dir = f"{output_dir}/final" model.save(final_output_dir) print(f"Model saved to {final_output_dir}") # ───────────────────────────────────────────────────────────── # 8. Benchmark test evaluation (commented out) # ───────────────────────────────────────────────────────────── # 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)