HachimiMT-demo / src /token_chunker.py
ngocdang83's picture
fix(chunk): paragraph mode gom dong - token_chunker.py
63a89c8 verified
Raw
History Blame
4.04 kB
"""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":
# "Theo đoạn": GOM cả đoạn (các dòng nối lại) thành một khối rồi mới chia,
# để sentence_chunks pack tới sát cap — KHÁC "sentence" (xử lý từng dòng
# riêng). Đoạn ngăn nhau bằng dòng trống vẫn tách (không trộn 2 đoạn).
paragraphs = [p.strip() for p in re.split(r"\n\s*\n+", text) if p.strip()]
chunks: list[str] = []
for paragraph in paragraphs:
lines = [line.strip() for line in paragraph.splitlines() if line.strip()]
if not lines:
continue
block = "\n".join(lines)
chunks.extend(sentence_chunks(tokenizer, block, 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