File size: 10,966 Bytes
33f617f
 
 
 
 
 
 
 
 
7812b29
33f617f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
import json
import random
import re
import hashlib
import unicodedata
from pathlib import Path
import numpy as np
import torch
import sentencepiece as spm
from .utils import log_progress

project_root = Path(__file__).resolve().parent.parent.parent
TOKENIZER_FILE = project_root / "src" / "tokenizer" / "tokenizer.model"

# SentencePiece của SAI không bật pad_id; id 0 là [UNK]. Model luôn nhận
# attention_mask tường minh nên dùng 0 để đệm vẫn an toàn.
PAD_ID = 0

_WHITESPACE = re.compile(r"\s+")

def normalize_text(text: str) -> str:
    """NFC + lowercase + gộp khoảng trắng, khớp với dữ liệu tokenizer đã học.

    Tokenizer được train trên text NFC đã lowercase: cùng một câu ở dạng NFD
    (dấu tách rời) bị cắt thành số token gấp ~4 lần và gần như là OOV.
    """
    text = unicodedata.normalize("NFC", text)
    return _WHITESPACE.sub(" ", text).strip().lower()

def text_key(text: str) -> int:
    """Hash ổn định của text đã chuẩn hoá, dùng để phát hiện trùng lặp / false negative."""
    return int(hashlib.md5(normalize_text(text).encode("utf-8")).hexdigest()[:15], 16)

def resolve_path(path) -> Path:
    path = Path(path)
    return path if path.is_absolute() else project_root / path

class EmbeddingTokenizer:
    def __init__(self, model_file: Path = TOKENIZER_FILE):
        self.model_file = str(model_file)
        self.sp = spm.SentencePieceProcessor()
        self.sp.load(self.model_file)
        self.bos_id = self.sp.piece_to_id("[BOS]")
        self.eos_id = self.sp.piece_to_id("[EOS]")

    # Cho phép pickle sang DataLoader worker (macOS dùng spawn).
    def __getstate__(self):
        return {"model_file": self.model_file}

    def __setstate__(self, state):
        self.__init__(state["model_file"])

    def encode(self, texts, prefix: str = "", max_len: int = 512):
        """[BOS] + tokens(prefix + text) + [EOS], cắt đuôi text nếu vượt max_len."""
        normalized = [normalize_text(prefix + t) for t in texts]
        pieces = self.sp.encode(normalized, out_type=int)
        return [[self.bos_id] + ids[:max_len - 2] + [self.eos_id] for ids in pieces]

    @staticmethod
    def pad(sequences, multiple_of: int = 8):
        """Right-padding. Trả về (input_ids, attention_mask, has_padding)."""
        max_len = max(len(s) for s in sequences)
        max_len = ((max_len + multiple_of - 1) // multiple_of) * multiple_of
        input_ids = torch.full((len(sequences), max_len), PAD_ID, dtype=torch.long)
        attention_mask = torch.zeros((len(sequences), max_len), dtype=torch.long)
        for i, s in enumerate(sequences):
            input_ids[i, :len(s)] = torch.as_tensor(s, dtype=torch.long)
            attention_mask[i, :len(s)] = 1
        has_padding = any(len(s) < max_len for s in sequences)
        return input_ids, attention_mask, has_padding

    def batch(self, texts, prefix: str = "", max_len: int = 512):
        return self.pad(self.encode(texts, prefix, max_len))

def iter_jsonl(path):
    with open(path, "r", encoding="utf-8") as f:
        for line in f:
            line = line.strip()
            if not line:
                continue
            try:
                yield json.loads(line)
            except json.JSONDecodeError:
                continue

def write_jsonl(path, rows) -> int:
    path = Path(path)
    path.parent.mkdir(parents=True, exist_ok=True)
    count = 0
    with open(path, "w", encoding="utf-8") as f:
        for row in rows:
            f.write(json.dumps(row, ensure_ascii=False) + "\n")
            count += 1
    return count

class PairDataset(torch.utils.data.Dataset):
    """Đọc nhiều file JSONL {query, positive, negatives?, type?} theo byte offset.

    Chỉ giữ offset trong RAM (không nạp toàn bộ text), mỗi worker tự mở file.
    `type == "symmetric"` (vd. NLI) nghĩa là hai vế cùng loại -> dùng query prefix cho cả hai.
    """
    def __init__(self, files: dict, max_rows_per_source: dict = None, seed: int = 54):
        max_rows_per_source = max_rows_per_source or {}
        rng = np.random.default_rng(seed)
        self.paths, self.sources = [], []
        offsets, file_idx = [], []

        for source, path in files.items():
            path = resolve_path(path)
            if not path.exists():
                log_progress(f"[data] Bỏ qua '{source}': không thấy {path}")
                continue
            source_offsets = []
            with open(path, "rb") as f:
                while True:
                    pos = f.tell()
                    line = f.readline()
                    if not line:
                        break
                    if line.strip():
                        source_offsets.append(pos)
            source_offsets = np.asarray(source_offsets, dtype=np.int64)
            cap = max_rows_per_source.get(source)
            if cap is not None and len(source_offsets) > cap:
                source_offsets = np.sort(rng.choice(source_offsets, size=cap, replace=False))
            self.paths.append(path)
            self.sources.append(source)
            offsets.append(source_offsets)
            file_idx.append(np.full(len(source_offsets), len(self.paths) - 1, dtype=np.int16))
            log_progress(f"[data] {source}: {len(source_offsets):,} mẫu")

        if not offsets:
            raise FileNotFoundError("Không có file train nào tồn tại")
        self.offsets = np.concatenate(offsets)
        self.source_idx = np.concatenate(file_idx)
        self._handles = {}

    def __len__(self):
        return len(self.offsets)

    def __getstate__(self):
        state = self.__dict__.copy()
        state["_handles"] = {}
        return state

    def __getitem__(self, idx):
        file_id = int(self.source_idx[idx])
        handle = self._handles.get(file_id)
        if handle is None:
            handle = open(self.paths[file_id], "rb")
            self._handles[file_id] = handle
        handle.seek(int(self.offsets[idx]))
        row = json.loads(handle.readline().decode("utf-8"))
        row["source"] = self.sources[file_id]
        return row

    def source_counts(self):
        counts = np.bincount(self.source_idx, minlength=len(self.sources))
        return dict(zip(self.sources, counts.tolist()))

class SourceBatchSampler(torch.utils.data.Sampler):
    """Mỗi batch chỉ chứa mẫu của MỘT nguồn.

    In-batch negative cùng domain khó hơn (luật với luật, tin tức với tin tức) và
    không trộn task đối xứng (NLI) với bất đối xứng (query -> passage) trong một loss.
    Thứ tự batch giữa các nguồn được xáo, nên tỷ lệ nguồn tỷ lệ với số mẫu.
    `repeats` {chỉ số nguồn: r} cho nguồn nhỏ (vd. dữ liệu chuyên ngành) đi r lượt mỗi epoch; mỗi lượt
    xáo riêng nên một batch không bao giờ chứa hai bản của cùng một dòng.
    """
    def __init__(self, source_idx, batch_size: int, seed: int = 54, repeats: dict = None):
        self.source_idx = np.asarray(source_idx)
        self.batch_size = batch_size
        self.seed = seed
        self.repeats = repeats or {}
        self.epoch = 0
        self.start_batch = 0

    def set_epoch(self, epoch: int, start_batch: int = 0):
        self.epoch = epoch
        self.start_batch = start_batch

    def _batches(self):
        rng = np.random.default_rng(self.seed + self.epoch)
        batches = []
        for source in np.unique(self.source_idx):
            members = np.flatnonzero(self.source_idx == source)
            for _ in range(self.repeats.get(int(source), 1)):
                idx = rng.permutation(members)
                n_full = len(idx) // self.batch_size
                if n_full == 0:
                    # Nguồn nhỏ hơn 1 batch vẫn được giữ lại thành một batch nhỏ.
                    batches.append(idx.tolist())
                    continue
                for b in range(n_full):
                    batches.append(idx[b * self.batch_size:(b + 1) * self.batch_size].tolist())
        order = rng.permutation(len(batches))
        return [batches[i] for i in order]

    def __len__(self):
        return len(self._batches())

    def batches_per_source(self, names):
        counts = np.bincount([self.source_idx[b[0]] for b in self._batches()], minlength=len(names))
        return dict(zip(names, counts.tolist()))

    def __iter__(self):
        yield from self._batches()[self.start_batch:]

class PairCollator:
    """`max_passage_len_by_source` ghi đè độ dài passage cho từng nguồn (vd. chunk tài liệu dài ~335
    token, cắt ở 256 thì mất đoạn chứa đáp án)."""
    def __init__(self, tokenizer: EmbeddingTokenizer, query_prefix: str, passage_prefix: str,
                 max_query_len: int, max_passage_len: int, num_negatives: int,
                 max_passage_len_by_source: dict = None):
        self.tokenizer = tokenizer
        self.query_prefix = query_prefix
        self.passage_prefix = passage_prefix
        self.max_query_len = max_query_len
        self.max_passage_len = max_passage_len
        self.max_passage_len_by_source = max_passage_len_by_source or {}
        self.num_negatives = num_negatives

    def __call__(self, rows):
        symmetric = rows[0].get("type") == "symmetric"
        passage_prefix = self.query_prefix if symmetric else self.passage_prefix
        max_passage_len = self.max_passage_len_by_source.get(rows[0]["source"], self.max_passage_len)

        queries = [r["query"] for r in rows]
        positives = [r["positive"] for r in rows]
        negatives = []
        for r in rows:
            negs = [n for n in (r.get("negatives") or []) if n]
            if len(negs) > self.num_negatives:
                negs = random.sample(negs, self.num_negatives)
            negatives.extend(negs)
        # passages = [positive của từng query] + [toàn bộ hard negative của batch]
        passages = positives + negatives

        q_ids, q_mask, q_pad = self.tokenizer.batch(queries, self.query_prefix, self.max_query_len)
        p_ids, p_mask, p_pad = self.tokenizer.batch(passages, passage_prefix, max_passage_len)
        return {
            "q_ids": q_ids, "q_mask": q_mask, "q_has_padding": q_pad,
            "p_ids": p_ids, "p_mask": p_mask, "p_has_padding": p_pad,
            "q_keys": torch.tensor([text_key(t) for t in queries], dtype=torch.long),
            "p_keys": torch.tensor([text_key(t) for t in passages], dtype=torch.long),
            # Nhóm của positive (trường "group", vd. mục tài liệu); negative không rõ nhóm -> -1.
            "p_groups": torch.tensor([text_key(f"group:{r['group']}") if r.get("group") else -1 for r in rows]
                                     + [-1] * len(negatives), dtype=torch.long),
            "source": rows[0]["source"],
        }