File size: 5,893 Bytes
58258b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
"""Tests for template auto-detection and stripping (no external deps)."""

from __future__ import annotations

from sdg.preprocessing.dedupe.templates import (
    find_common_prefixes,
    strip_template,
    template_hit_stats,
)


# ── find_common_prefixes ─────────────────────────────────────────────────────

def test_find_common_prefixes_empty_input():
    assert find_common_prefixes([], min_count=2) == []


def test_find_common_prefixes_returns_empty_when_no_common_prefix():
    texts = ["alpha question one", "beta question two", "gamma question three"]
    assert find_common_prefixes(texts, min_count=2, prefix_lengths=(5, 10)) == []


def test_find_common_prefixes_detects_simple_repeated_prefix():
    template = "Answer the question: "
    texts = [template + f"unique body number {i}" for i in range(20)]
    found = find_common_prefixes(texts, min_count=10, prefix_lengths=(len(template),))
    assert template in found


def test_find_common_prefixes_respects_min_count():
    rare = "Rare prefix: "
    common = "Common prefix: "
    texts = (
        [rare + f"a{i}" for i in range(5)]
        + [common + f"b{i}" for i in range(20)]
    )
    found = find_common_prefixes(texts, min_count=10, prefix_lengths=(len(common),))
    assert common in found
    assert rare not in found


def test_find_common_prefixes_returns_sorted_longest_first():
    long_tmpl = "A" * 100 + " specific"
    short_tmpl = "A" * 50
    texts = [long_tmpl + " body " + str(i) for i in range(20)]
    found = find_common_prefixes(texts, min_count=10, prefix_lengths=(50, 100, 150))
    assert found == sorted(found, key=len, reverse=True)


def test_find_common_prefixes_handles_short_texts():
    """Texts shorter than a probed prefix length must not crash."""
    texts = ["xy"] * 20  # all shorter than every probed length
    found = find_common_prefixes(texts, min_count=2, prefix_lengths=(50, 100))
    assert found == []


def test_find_common_prefixes_finds_multiple_distinct_templates():
    t1 = "Template A: "
    t2 = "Different template B: "
    texts = (
        [t1 + f"x{i}" for i in range(15)]
        + [t2 + f"y{i}" for i in range(15)]
        + [f"singleton {i}" for i in range(5)]
    )
    found = find_common_prefixes(texts, min_count=10, prefix_lengths=(len(t1), len(t2)))
    assert t1 in found
    assert t2 in found


# ── strip_template ───────────────────────────────────────────────────────────

def test_strip_template_returns_text_unchanged_with_no_match():
    assert strip_template("hello world", ["FOO: "]) == "hello world"


def test_strip_template_returns_text_unchanged_with_empty_templates():
    assert strip_template("hello world", []) == "hello world"


def test_strip_template_simple_match():
    assert strip_template("PREFIX: actual content", ["PREFIX: "]) == "actual content"


def test_strip_template_strips_leading_whitespace_after_removal():
    assert strip_template("PREFIX:    \n\ncontent", ["PREFIX:"]) == "content"


def test_strip_template_longest_match_wins_when_sorted_correctly():
    """Templates list sorted longest-first as find_common_prefixes returns."""
    templates = ["PREFIX EXTRA: ", "PREFIX: "]
    assert strip_template("PREFIX EXTRA: content", templates) == "content"
    assert strip_template("PREFIX: content", templates) == "content"


def test_strip_template_does_not_match_substring_in_middle():
    """Template must occur as a *prefix*, not anywhere in the text."""
    assert strip_template("hello PREFIX: world", ["PREFIX: "]) == "hello PREFIX: world"


# ── template_hit_stats ───────────────────────────────────────────────────────

def test_template_hit_stats_partitions_texts_correctly():
    texts = [
        "MCQ: question 1",
        "MCQ: question 2",
        "MATH: question 3",
        "totally distinct prompt",
    ]
    templates = ["MCQ: ", "MATH: "]
    stats = template_hit_stats(texts, templates)
    assert stats["MCQ: "] == 2
    assert stats["MATH: "] == 1
    assert stats["<no template>"] == 1
    assert sum(stats.values()) == len(texts)


def test_template_hit_stats_uses_longest_match_when_templates_nest():
    """If templates list has nested prefixes, longest-first prevents double counting."""
    texts = ["PREFIX EXTRA: x"] * 10 + ["PREFIX: y"] * 5
    templates = ["PREFIX EXTRA: ", "PREFIX: "]
    stats = template_hit_stats(texts, templates)
    assert stats["PREFIX EXTRA: "] == 10
    assert stats["PREFIX: "] == 5
    assert stats["<no template>"] == 0


def test_template_hit_stats_empty_templates_all_unmatched():
    texts = ["a", "b", "c"]
    stats = template_hit_stats(texts, [])
    assert stats == {"<no template>": 3}


# ── End-to-end: detect then strip ────────────────────────────────────────────

def test_detect_then_strip_round_trip():
    """Realistic sequence: detect templates, then strip them from the corpus."""
    template = "Answer the following multiple choice question: "
    questions = [
        "What is 2+2?",
        "What color is the sky?",
        "Who wrote Hamlet?",
        "What is the capital of France?",
        "What is the speed of light?",
    ] * 4  # 20 questions
    texts = [template + q for q in questions]
    templates = find_common_prefixes(texts, min_count=10, prefix_lengths=(len(template),))
    stripped = [strip_template(t, templates) for t in texts]
    # Every text should have been stripped down to just its question.
    assert all(s in [q for q in questions] for s in stripped)