UCAS-EasyTranslate / src /easytranslate /data /preprocessing.py
sdfjliom's picture
Implement data pipeline module (#2)
6b73a07
Raw
History Blame Contribute Delete
3.91 kB
"""Text cleaning and filtering utilities for translation corpora."""
from __future__ import annotations
import re
import unicodedata
from typing import Callable, Iterable, Sequence
CONTROL_OR_ZERO_WIDTH_RE = re.compile(r"[\u0000-\u001f\u007f-\u009f\u200b\u200c\u200d\ufeff]")
SPACE_RE = re.compile(r"\s+")
PUNCT_RE = re.compile(r"[^\w\s\u4e00-\u9fff]", flags=re.UNICODE)
def clean_text(text: str, lowercase: bool = False, remove_punctuation: bool = False) -> str:
"""Normalize a single sentence without changing its meaning aggressively."""
if text is None:
return ""
text = unicodedata.normalize("NFKC", str(text))
text = CONTROL_OR_ZERO_WIDTH_RE.sub("", text)
text = SPACE_RE.sub(" ", text).strip()
if lowercase:
text = text.lower()
if remove_punctuation:
text = PUNCT_RE.sub("", text)
text = SPACE_RE.sub(" ", text).strip()
return text
def _default_length(text: str) -> int:
"""Use whitespace tokens for Latin text and character count for CJK-heavy text."""
cjk_chars = sum(1 for ch in text if "\u4e00" <= ch <= "\u9fff")
if cjk_chars >= max(1, len(text) // 3):
return len(text.replace(" ", ""))
return len(text.split())
def filter_by_length(
src: str,
tgt: str,
min_src_len: int = 1,
min_tgt_len: int = 1,
max_src_len: int = 256,
max_tgt_len: int = 256,
length_ratio_threshold: float = 3.0,
length_fn: Callable[[str], int] | None = None,
) -> bool:
"""Return True when a sentence pair passes basic length and ratio checks."""
length_fn = length_fn or _default_length
src_len = length_fn(src)
tgt_len = length_fn(tgt)
if src_len < min_src_len or tgt_len < min_tgt_len:
return False
if src_len > max_src_len or tgt_len > max_tgt_len:
return False
shorter = max(1, min(src_len, tgt_len))
longer = max(src_len, tgt_len)
return longer / shorter <= length_ratio_threshold
def deduplicate_pairs(pairs: Iterable[tuple[str, str]]) -> list[tuple[str, str]]:
"""Deduplicate by exact cleaned source-target pair while preserving order."""
seen: set[tuple[str, str]] = set()
result: list[tuple[str, str]] = []
for src, tgt in pairs:
key = (src, tgt)
if key in seen:
continue
seen.add(key)
result.append(key)
return result
def preprocess_pipeline(
src_texts: Sequence[str],
tgt_texts: Sequence[str],
lowercase_src: bool = False,
lowercase_tgt: bool = False,
remove_punctuation: bool = False,
max_src_len: int = 256,
max_tgt_len: int = 256,
min_src_len: int = 1,
min_tgt_len: int = 1,
filter_by_length_enabled: bool = True,
length_ratio_threshold: float = 3.0,
deduplicate: bool = True,
) -> tuple[list[str], list[str]]:
"""Clean, filter, and optionally deduplicate parallel source-target texts."""
if len(src_texts) != len(tgt_texts):
raise ValueError("src_texts and tgt_texts must have the same length")
pairs: list[tuple[str, str]] = []
for raw_src, raw_tgt in zip(src_texts, tgt_texts):
src = clean_text(raw_src, lowercase=lowercase_src, remove_punctuation=remove_punctuation)
tgt = clean_text(raw_tgt, lowercase=lowercase_tgt, remove_punctuation=remove_punctuation)
if not src or not tgt:
continue
if filter_by_length_enabled and not filter_by_length(
src,
tgt,
min_src_len=min_src_len,
min_tgt_len=min_tgt_len,
max_src_len=max_src_len,
max_tgt_len=max_tgt_len,
length_ratio_threshold=length_ratio_threshold,
):
continue
pairs.append((src, tgt))
if deduplicate:
pairs = deduplicate_pairs(pairs)
if not pairs:
return [], []
src_clean, tgt_clean = zip(*pairs)
return list(src_clean), list(tgt_clean)