File size: 3,655 Bytes
1a212f3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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)