File size: 3,154 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
96
97
98
99
100
101
102
103
104
105
"""
Fine-tune a cross-encoder reranker for Turkish legal QA.

Training data (weak supervision):
- Positive: (Soru, Cevap) from HF train- Negative: (Soru, Cevap') where Cevap' is a random other answer

Output:
  A folder loadable by sentence_transformers.CrossEncoder
"""

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
from sentence_transformers.cross_encoder import CrossEncoder
from sentence_transformers.cross_encoder.evaluation import CEBinaryClassificationEvaluator
from torch.utils.data import DataLoader


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

    qs: List[str] = []
    ps: 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
        qs.append(q)
        ps.append(p)

    n = len(qs)
    if n < 2:
        raise RuntimeError("Not enough HF train rows.")

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

    train_examples: List[InputExample] = []
    eval_examples: List[InputExample] = []

    for i, idx in enumerate(indices):
        q = qs[idx]
        pos = ps[idx]
        neg_j = random.randint(0, n - 1)
        if neg_j == idx:
            neg_j = (neg_j + 1) % n
        neg = ps[neg_j]

        ex_pos = InputExample(texts=[q, pos], label=1.0)
        ex_neg = InputExample(texts=[q, neg], label=0.0)

        # simple 90/10 split by modulo
        if i % 10 == 0:
            eval_examples.extend([ex_pos, ex_neg])
        else:
            train_examples.extend([ex_pos, ex_neg])

    return train_examples, eval_examples


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--base_model", default="cross-encoder/ms-marco-MiniLM-L-6-v2")
    parser.add_argument("--out_dir", default="models/ce-legal-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=6000)
    parser.add_argument("--seed", type=int, default=42)
    args = parser.parse_args()

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

    train_ex, eval_ex = build_pairs(args.max_samples, args.seed)
    train_loader = DataLoader(train_ex, shuffle=True, batch_size=args.batch_size)

    model = CrossEncoder(args.base_model, num_labels=1)
    evaluator = CEBinaryClassificationEvaluator.from_input_examples(eval_ex, name="legal-dev")

    model.fit(
        train_dataloader=train_loader,
        evaluator=evaluator,
        epochs=args.epochs,
        warmup_steps=max(100, int(len(train_loader) * 0.1)),
        output_path=str(out),
        show_progress_bar=True,
    )

    print(f"[RerankerFT] Saved cross-encoder -> {out}")


if __name__ == "__main__":
    main()