File size: 2,814 Bytes
6dfa658
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Contrastive fine-tuning for dense retrieval (sentence-transformers).

Pairs:
  (query=Soru, passage=Cevap) from HF ``train`` split

Hard negatives (simple):
  For each query, take one random other row's Cevap as a negative passage.

Output:
  A local folder suitable for SentenceTransformer loading, e.g. models/st-legal-v1
"""

from __future__ import annotations

import argparse
import random
from pathlib import Path
from typing import List, Tuple

from datasets import load_dataset
from sentence_transformers import InputExample, SentenceTransformer, losses
from torch.utils.data import DataLoader


def build_examples(max_samples: int, seed: int) -> List[InputExample]:
    random.seed(seed)
    ds = load_dataset("Renicames/turkish-law-chatbot")
    train = ds["train"]

    texts_q: List[str] = []
    texts_p: List[str] = []
    for row in train:
        q = str(row.get("Soru", "")).strip()
        p = str(row.get("Cevap", "")).strip()
        if not q or not p:
            continue
        texts_q.append(q)
        texts_p.append(p)

    n = len(texts_q)
    if n == 0:
        raise RuntimeError("No usable HF train rows.")

    limit = min(max_samples, n)
    idxs = list(range(n))
    random.shuffle(idxs)
    idxs = idxs[:limit]

    examples: List[InputExample] = []
    for i in idxs:
        q = texts_q[i]
        pos = texts_p[i]
        neg_j = random.randint(0, n - 1)
        if neg_j == i:
            neg_j = (neg_j + 1) % n
        neg = texts_p[neg_j]
        examples.append(InputExample(texts=[q, pos, neg]))
    return examples


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--base_model", default="sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2")
    parser.add_argument("--out_dir", default="models/st-legal-multilingual-v1")
    parser.add_argument("--epochs", type=int, default=1)
    parser.add_argument("--batch_size", type=int, default=16)
    parser.add_argument("--max_samples", type=int, default=8000)
    parser.add_argument("--warmup_ratio", type=float, default=0.05)
    parser.add_argument("--seed", type=int, default=42)
    args = parser.parse_args()

    out = Path(args.out_dir)
    out.mkdir(parents=True, exist_ok=True)

    examples = build_examples(args.max_samples, args.seed)
    train_dataloader = DataLoader(examples, shuffle=True, batch_size=args.batch_size)

    model = SentenceTransformer(args.base_model)
    train_loss = losses.MultipleNegativesRankingLoss(model)

    model.fit(
        train_objectives=[(train_dataloader, train_loss)],
        epochs=args.epochs,
        warmup_steps=int(len(train_dataloader) * args.warmup_ratio),
        show_progress_bar=True,
        output_path=str(out),
    )

    print(f"[EmbeddingFT] Saved model -> {out}")


if __name__ == "__main__":
    main()