File size: 4,049 Bytes
ba3ecf1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
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)