jstAnotherCapi's picture
Upload folder using huggingface_hub
ba3ecf1 verified
Raw
History Blame Contribute Delete
4.05 kB
import os
import sys
import torch
import logging
from datetime import datetime
from datasets import load_dataset
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
# --- MANUAL CHECKPOINT CONFIG ---
# PASTE YOUR CHECKPOINT PATH HERE to resume. Example: "output/arabert_20240520_1530/checkpoint-6000"
# Set to None if you want to start a brand new training run.
CHECKPOINT_PATH = "/home/skiredj.abderrahman/khalil/sbert_training/output/arabert_20260224_1730/checkpoint-102000"
model_name = "bert-base-arabertv02"
model_nickname = "arabert"
train_batch_size = 64
# Logic to maintain the same output directory if resuming
if CHECKPOINT_PATH and os.path.exists(CHECKPOINT_PATH):
output_dir = str(os.path.dirname(CHECKPOINT_PATH))
print(f"--- RESUMING FROM: {CHECKPOINT_PATH} ---")
else:
timestamp = datetime.now().strftime("%Y%m%d_%H%M")
output_dir = f"output/{model_nickname}_{timestamp}"
CHECKPOINT_PATH = None
print(f"--- STARTING NEW RUN: {output_dir} ---")
# --- 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"))
# --- MODEL & DATA ---
device = "cuda" if torch.cuda.is_available() else "cpu"
model = SentenceTransformer(model_name, device=device)
train_dataset = load_dataset("csv", data_files="train.csv")
eval_dataset = load_dataset("csv", data_files="val.csv")
test_dataset = load_dataset("csv", data_files="test.csv")
eval_subset = eval_dataset["train"].shuffle(seed=42).select(range(min(250000, len(eval_dataset["train"]))))
# --- LOSS & EVALUATORS ---
matryoshka_dims = [768, 512, 256, 128, 64]
inner_train_loss = losses.MultipleNegativesRankingLoss(model=model)
train_loss = losses.MatryoshkaLoss(model, inner_train_loss, matryoshka_dims=matryoshka_dims)
evaluators = [
TripletEvaluator(
anchors=eval_subset["anchor"],
positives=eval_subset["positive"],
negatives=eval_subset["negative"],
name=f"dev-{dim}",
truncate_dim=dim,
) for dim in matryoshka_dims
]
dev_evaluator = SequentialEvaluator(evaluators, main_score_function=lambda scores: scores[0])
# --- TRAINING ARGS ---
args = SentenceTransformerTrainingArguments(
output_dir=output_dir,
num_train_epochs=4,
per_device_train_batch_size=train_batch_size,
gradient_accumulation_steps=2,
bf16=True,
learning_rate=2e-5,
warmup_ratio=0.1,
batch_sampler=BatchSamplers.NO_DUPLICATES,
eval_strategy="steps",
eval_steps=6000,
save_strategy="steps",
save_steps=6000,
save_total_limit=2,
logging_steps=200,
)
trainer = SentenceTransformerTrainer(
model=model,
args=args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
loss=train_loss,
evaluator=dev_evaluator,
)
# --- THE TRAIN CALL ---
trainer.train(resume_from_checkpoint=CHECKPOINT_PATH)
# Save final model
final_output_dir = "/home/skiredj.abderrahman/khalil/sbert_training/output/final_epoch4"
model.save(final_output_dir)
print("model saved successfully")
# Test evaluation
evaluators = []
for dim in matryoshka_dims:
evaluators.append(
TripletEvaluator(
anchors=test_dataset["train"]["anchor"],
positives=test_dataset["train"]["positive"],
negatives=test_dataset["train"]["negative"],
name=f"test-{dim}",
truncate_dim=dim,
)
)
test_evaluator = SequentialEvaluator(evaluators)
test_evaluator(model)