"""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)