Download sdg/preprocessing/dedupe/test_templates.py from fzzhang/svd-code: direct link, hf CLI and curl.
- Browser
- Download file 5.89 kB
-
https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/preprocessing/dedupe/test_templates.py
- Command line
-
hf download hf://fzzhang/svd-code/sdg/preprocessing/dedupe/test_templates.py
-
curl -L -o test_templates.py https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/preprocessing/dedupe/test_templates.py
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) | |