File size: 3,707 Bytes
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
"""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

SENTENCE_RE = re.compile(r"[^。!?!?;;]+[。!?!?;;]*")


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 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":
        paragraphs = [p.strip() for p in re.split(r"\n\s*\n+", text) if p.strip()]
        chunks: list[str] = []
        for paragraph in paragraphs:
            for line in paragraph.splitlines():
                line = line.strip()
                if not line:
                    continue
                chunks.extend(sentence_chunks(tokenizer, line, 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