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