File size: 7,869 Bytes
e9015b1
 
 
 
 
 
 
 
 
 
9f454ca
 
 
 
 
 
 
56f65cd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e9015b1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
977b7ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3ba49c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e9015b1
 
 
 
 
 
 
 
 
 
 
 
 
3ba49c2
 
 
 
 
e9015b1
 
 
 
 
 
 
 
 
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
"""Token-aware chunking for Marian models (ported from HachimiMT HF Space)."""

from __future__ import annotations

import re
from typing import TYPE_CHECKING

if TYPE_CHECKING:
    from transformers import PreTrainedTokenizerBase

# Câu nguồn: nuốt tới cụm dấu kết (gồm 。!?!?;; và … ……) kèm ngoặc đóng theo
# sau (」』)》】 và ”’ " '), hoặc hết chuỗi. Giữ ngoặc đóng Ở LẠI câu vừa kết,
# không đẩy sang câu sau (trước đây 」 sau 。 rơi nhầm vào câu kế tiếp).
SENTENCE_RE = re.compile(
    r".+?(?:[。!?!?;;…]+[」』)》】”’\"']*|$)",
    re.S,
)
HEADING_RE = re.compile(r"^第[0-9零〇一二三四五六七八九十百千万两]+[章节回卷部篇]")
METADATA_RE = re.compile(r"^(?:书名|作者|简介|内容简介|作品简介)\s*[::]")


def is_hard_boundary_line(line: str) -> bool:
    """Lines that must stay standalone even in paragraph chunk mode."""
    stripped = (line or "").strip()
    if not stripped or "\n" in stripped:
        return False
    if METADATA_RE.match(stripped):
        return True
    if len(stripped) > 32:
        return False
    if re.search(r"[。!?!?;;,,::\"“”]", stripped):
        return False
    return bool(HEADING_RE.match(stripped))


def source_token_ids(
    tokenizer: PreTrainedTokenizerBase,
    text: str,
    *,
    max_length: int,
    truncation: bool,
) -> list[int]:
    token_ids = tokenizer(
        text,
        truncation=truncation,
        max_length=max_length,
    )["input_ids"]
    if tokenizer.pad_token_id is not None:
        token_ids = [tid for tid in token_ids if tid != tokenizer.pad_token_id]
    return token_ids


def source_token_count(
    tokenizer: PreTrainedTokenizerBase,
    text: str,
    *,
    max_length: int,
) -> int:
    return len(source_token_ids(tokenizer, text, max_length=max_length, truncation=False))


def char_chunks(
    tokenizer: PreTrainedTokenizerBase,
    text: str,
    *,
    max_tokens: int,
) -> list[str]:
    chunks: list[str] = []
    remaining = text
    while remaining:
        if source_token_count(tokenizer, remaining, max_length=max_tokens) <= max_tokens:
            chunks.append(remaining)
            break

        low, high = 1, len(remaining)
        best = 1
        while low <= high:
            middle = (low + high) // 2
            candidate = remaining[:middle]
            if source_token_count(tokenizer, candidate, max_length=max_tokens) <= max_tokens:
                best = middle
                low = middle + 1
            else:
                high = middle - 1

        chunks.append(remaining[:best])
        remaining = remaining[best:]
    return chunks


def sentence_chunks(
    tokenizer: PreTrainedTokenizerBase,
    line: str,
    *,
    max_tokens: int,
) -> list[str]:
    if source_token_count(tokenizer, line, max_length=max_tokens) <= max_tokens:
        return [line]

    pieces = [match.group(0) for match in SENTENCE_RE.finditer(line)]
    if not pieces:
        return char_chunks(tokenizer, line, max_tokens=max_tokens)

    chunks: list[str] = []
    current = ""
    for piece in pieces:
        if source_token_count(tokenizer, piece, max_length=max_tokens) > max_tokens:
            if current:
                chunks.append(current)
                current = ""
            chunks.extend(char_chunks(tokenizer, piece, max_tokens=max_tokens))
            continue

        candidate = current + piece
        if current and source_token_count(tokenizer, candidate, max_length=max_tokens) > max_tokens:
            chunks.append(current)
            current = piece
        else:
            current = candidate

    if current:
        chunks.append(current)
    return chunks


def _layout_lines(text: str) -> list[str]:
    """Tách dòng GIỮ dòng rỗng đầu/cuối, chuẩn hóa newline (cho plan-based)."""
    return (text or "").replace("\r\n", "\n").replace("\r", "\n").split("\n")


def split_sentence_lines_with_plan(
    tokenizer: PreTrainedTokenizerBase,
    text: str,
    *,
    max_tokens: int,
) -> tuple[list[str], list[list[int] | None]]:
    """Chia "Theo câu" KÈM plan ánh xạ dòng-nguồn → chunk-index.

    Trả `(chunks, plan)` với `plan` theo TỪNG dòng nguồn (giữ cả dòng trống):
    `None` nếu dòng trống, hoặc danh sách index các chunk thuộc dòng đó (một dòng
    dài vượt cap có thể tách thành nhiều chunk → nhiều index). Cho phép ghép lại
    output theo đúng bố cục dòng nguồn (giữ dòng trống), điều mà `split_for_
    translation` thường (trả phẳng list[str]) không làm được.
    """
    chunks: list[str] = []
    plan: list[list[int] | None] = []
    for line in _layout_lines(text):
        stripped = line.strip()
        if not stripped:
            plan.append(None)
            continue
        line_chunks = sentence_chunks(tokenizer, stripped, max_tokens=max_tokens)
        indices = list(range(len(chunks), len(chunks) + len(line_chunks)))
        chunks.extend(line_chunks)
        plan.append(indices)
    return chunks, plan


def split_paragraphs_with_plan(
    tokenizer: PreTrainedTokenizerBase,
    text: str,
    *,
    max_tokens: int,
) -> tuple[list[str], list[tuple[int, ...]]]:
    """Chia "Theo đoạn" thật, kèm map chunk -> line-index tuyệt đối.

    Khác "Theo câu": các dòng liên tiếp trong cùng đoạn được pack vào một chunk
    tới sát cap để giảm số chunk và giữ context. Khác bản cũ: mỗi chunk giữ danh
    sách line-index nguồn, cho phép restore/fallback theo từng chunk thay vì đoán
    toàn tài liệu.
    """
    chunks: list[str] = []
    plan: list[tuple[int, ...]] = []
    buffered_lines: list[str] = []
    buffered_indices: list[int] = []

    def flush_buffer() -> None:
        if not buffered_lines:
            return
        chunks.append("\n".join(buffered_lines))
        plan.append(tuple(buffered_indices))
        buffered_lines.clear()
        buffered_indices.clear()

    def add_line_chunk(line_index: int, line: str) -> None:
        for piece in sentence_chunks(tokenizer, line, max_tokens=max_tokens):
            chunks.append(piece)
            plan.append((line_index,))

    for line_index, line in enumerate(_layout_lines(text)):
        stripped = line.strip()
        if not stripped:
            flush_buffer()
            continue

        if is_hard_boundary_line(stripped):
            flush_buffer()
            add_line_chunk(line_index, stripped)
            continue

        if source_token_count(tokenizer, stripped, max_length=max_tokens) > max_tokens:
            flush_buffer()
            add_line_chunk(line_index, stripped)
            continue

        candidate_lines = [*buffered_lines, stripped]
        candidate = "\n".join(candidate_lines)
        if buffered_lines and source_token_count(tokenizer, candidate, max_length=max_tokens) > max_tokens:
            flush_buffer()
        buffered_lines.append(stripped)
        buffered_indices.append(line_index)

    flush_buffer()
    return chunks, plan


def split_for_translation(
    tokenizer: PreTrainedTokenizerBase,
    text: str,
    *,
    max_tokens: int,
    chunk_mode: str = "sentence",
) -> list[str]:
    """Split *text* into chunks that fit within *max_tokens*."""
    text = text.strip()
    if not text:
        return []

    if chunk_mode == "paragraph":
        chunks, _plan = split_paragraphs_with_plan(
            tokenizer,
            text,
            max_tokens=max_tokens,
        )
        return chunks

    chunks = []
    for line in text.splitlines():
        line = line.strip()
        if not line:
            continue
        chunks.extend(sentence_chunks(tokenizer, line, max_tokens=max_tokens))
    return chunks