File size: 3,516 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
import os
import sys
import torch
import logging
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
import glob

# --- CONFIG ---
MODEL_PATH = "/home/skiredj.abderrahman/khalil/sbert_training/epoch2/model/"
train_batch_size = 64
output_dir = "output/arabert_ms_marco"

# --- LOGGING ---
logging.basicConfig(
    format="%(asctime)s - %(message)s",
    datefmt="%Y-%m-%d %H:%M:%S",
    level=logging.INFO,
    handlers=[logging.FileHandler("logs_ds36.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_ds36.txt", "a"))

# --- LOAD DATA ---
train_dataset = load_dataset("csv", data_files="clean_dataset36_train.csv")["train"]
val_dataset = load_dataset("csv", data_files="clean_dataset36_val.csv")["train"]
print(f"Train size: {len(train_dataset)} | Val size: {len(val_dataset)}")

# --- MODEL ---
device = "cuda" if torch.cuda.is_available() else "cpu"
model = SentenceTransformer(MODEL_PATH, device=device)

# --- 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=val_dataset["anchor"],
        positives=val_dataset["positive"],
        negatives=val_dataset["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=2,
    per_device_train_batch_size=train_batch_size,
    gradient_accumulation_steps=2,
    bf16=True,
    learning_rate=1e-5,
    warmup_ratio=0.1,
    batch_sampler=BatchSamplers.NO_DUPLICATES,
    eval_strategy="steps",
    eval_steps=12000,
    save_strategy="steps",
    save_steps=12000,
    save_total_limit=2,
    logging_steps=200,
)

# --- RESUME IF CHECKPOINT EXISTS ---
existing_checkpoints = sorted(glob.glob(f"{output_dir}/checkpoint-*"))
resume_from = existing_checkpoints[-1] if existing_checkpoints else None
if resume_from:
    print(f"Resuming from: {resume_from}")
else:
    print("Starting fresh")

# --- TRAIN ---
trainer = SentenceTransformerTrainer(
    model=model,
    args=args,
    train_dataset=train_dataset,
    eval_dataset=val_dataset,
    loss=train_loss,
    evaluator=dev_evaluator,
)
trainer.train(resume_from_checkpoint=resume_from)

# --- SAVE FINAL ---
final_output_dir = "/home/skiredj.abderrahman/khalil/sbert_training/output/final_ms_marco"
model.save(final_output_dir)
print(f"Model saved to {final_output_dir}")

# --- FINAL EVAL ---
test_evaluators = [
    TripletEvaluator(
        anchors=val_dataset["anchor"],
        positives=val_dataset["positive"],
        negatives=val_dataset["negative"],
        name=f"final-{dim}",
        truncate_dim=dim,
    ) for dim in matryoshka_dims
]
SequentialEvaluator(test_evaluators)(model)