import re from collections.abc import Callable from difflib import SequenceMatcher from tokenizers.pre_tokenizers import Whitespace ATTACHED_PUNCTUATION = r',.;:!?\(\)\[\]\{}\'"\-_/\\|@#\$%\^&\*\+=<>~`' _ATTACHED_PUNCTUATION_PATTERN = re.compile(rf"([{ATTACHED_PUNCTUATION}])") WHITESPACE_PRE_TOKENIZER = Whitespace() TokenizeFn = Callable[[str], tuple[str, ...]] def split_attached_punctuation(token: str) -> list[str]: return [part for part in _ATTACHED_PUNCTUATION_PATTERN.split(token) if part] def tokenize_line(line: str) -> tuple[str, ...]: normalized = ( line .replace(", ", " , ") .replace(". ", " . ") .replace(";", " ; ") .replace(":", " : ") ) tokens: list[str] = [] for token in normalized.split(): tokens.extend(split_attached_punctuation(token)) return tuple(tokens) def tokenize_whitespace(text: str) -> tuple[str, ...]: return tuple(token for token, _ in WHITESPACE_PRE_TOKENIZER.pre_tokenize_str(text)) def join_tokenized(text: str) -> str: return " ".join(tokenize_line(text)) def join_natural(text: str) -> str: tokens = tokenize_line(text) if not tokens: return "" result = tokens[0] for token in tokens[1:]: if token in ",.;:!?)]}": result += token elif result.endswith(",") and token.isdigit(): result += token elif result and result[-1] in "([{": result += token else: result += " " + token return result def extract_labels( original_text: str, edited_text: str, tokenize: TokenizeFn = tokenize_line, ) -> tuple[int, ...]: """ Compare original and edited text at tokenize() granularity. Default tokenize_line() aligns natural prose (demo) and SwissGov-style pre-tokenized input. For DSD, pass tokenize_whitespace so labels match the Whitespace pre-tokenizer used by gold labels and spans_from_labels. """ original_tokens = tokenize(original_text) edited_tokens = tokenize(edited_text) labels: list[int] = [] for tag, start_original, end_original, _, _ in SequenceMatcher( None, original_tokens, edited_tokens ).get_opcodes(): if tag == "equal": labels.extend(0 for _ in range(start_original, end_original)) elif tag in ("delete", "replace"): labels.extend(1 for _ in range(start_original, end_original)) return tuple(labels) def extract_edit_tooltips( original_text: str, edited_text: str | None, tokenize: TokenizeFn = tokenize_line, ) -> tuple[str, ...]: """ For each original token marked as edited, return a tooltip describing the diff. """ original_tokens = tokenize(original_text) if edited_text is None: return tuple("" for _ in original_tokens) edited_tokens = tokenize(edited_text) tooltips: list[str] = [""] * len(original_tokens) for tag, start_original, end_original, start_edited, end_edited in SequenceMatcher( None, original_tokens, edited_tokens ).get_opcodes(): if tag == "equal": continue original_span = " ".join(original_tokens[start_original:end_original]) if tag == "delete": tooltip = f'Deleted: "{original_span}"' elif tag == "replace": edited_span = " ".join(edited_tokens[start_edited:end_edited]) tooltip = f'Replaced "{original_span}" with "{edited_span}"' else: continue for index in range(start_original, end_original): tooltips[index] = tooltip return tuple(tooltips)