jstAnotherCapi's picture
Upload folder using huggingface_hub
ba3ecf1 verified
Raw
History Blame Contribute Delete
15.2 kB
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("Skipping 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,
)
print('###################resuming the training##########################')
trainer.train(resume_from_checkpoint="output/arabert_20260311_2304/checkpoint-204000")
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 = {
"query": ss_query_prompt,
"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)