| 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.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")) |
|
|
| |
| |
| |
| 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] |
|
|
| |
| |
| |
| 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 = False |
| TEST_SAMPLES = 100 |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| logging.info("Loading datasets...") |
|
|
| def load_csv(path): |
| ds = load_dataset("csv", data_files=path)["train"] |
| |
| 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) |
|
|
| |
| |
| |
| 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)), |
| } |
|
|
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| 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", |
| ) |
|
|
| |
| |
| |
| 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") |
|
|
| |
| |
| |
|
|
| final_output_dir = f"{output_dir}/final" |
| model.save(final_output_dir) |
| print(f"Model saved to {final_output_dir}") |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|