Spaces:
Running
Running
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
|