svd-code / sdg /preprocessing /dedupe /test_templates.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
5.89 kB
"""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)