svd-code / sdg /preprocessing /dedupe /templates.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
2.97 kB
"""Auto-detect and strip common instruction-template prefixes for dedup similarity.
Many SFT datasets (especially MCQ-heavy ones) use a small number of fixed
instruction templates as prompt prefixes, e.g.
"Answer the following multiple choice question. The last line of your
response should be in the following format: 'Answer: A/B/C/D' (e.g.
'Answer: A')."
These shared prefixes inflate Jaccard similarity (MinHash) and embedding
cosine similarity (semantic), causing distinct questions that share a
template to be falsely merged as duplicates.
This module provides:
- find_common_prefixes(): scan a corpus and auto-detect frequent prefixes
- strip_template(): remove the longest matching template from a text
- template_hit_stats(): audit helper showing per-template hit counts
Intended usage: pass the *stripped* prompts to MinHash/Semantic for the
similarity computation, but keep the *original* (unstripped) prompts in
the records that get persisted/uploaded.
"""
from __future__ import annotations
from collections import Counter
def find_common_prefixes(
texts: list[str],
min_count: int = 10,
prefix_lengths: tuple[int, ...] = (50, 100, 150, 200, 300, 400, 500, 600),
) -> list[str]:
"""Find character-prefix substrings that begin >= `min_count` texts.
Tries each length in `prefix_lengths` independently — this lets the
detector find both short and long template variants. The result is sorted
longest-first so that during stripping the most-specific (longest) match
wins (see `strip_template`).
"""
counter: Counter = Counter()
for t in texts:
for length in prefix_lengths:
if len(t) >= length:
counter[t[:length]] += 1
templates = [pfx for pfx, c in counter.items() if c >= min_count]
templates.sort(key=len, reverse=True)
return templates
def strip_template(text: str, templates: list[str]) -> str:
"""Return `text` with the longest matching template prefix removed.
`templates` should be sorted longest-first (as returned by
`find_common_prefixes`). Leading whitespace is stripped from the result.
Returns `text` unchanged when no template matches or `templates` is empty.
"""
for tmpl in templates:
if text.startswith(tmpl):
return text[len(tmpl):].lstrip()
return text
def template_hit_stats(texts: list[str], templates: list[str]) -> dict[str, int]:
"""For each template, count how many texts have it as the longest matching prefix.
Texts that don't match any template are counted under '<no template>'.
Sum of counts equals len(texts).
"""
counts: dict[str, int] = {tmpl: 0 for tmpl in templates}
counts["<no template>"] = 0
for t in texts:
for tmpl in templates:
if t.startswith(tmpl):
counts[tmpl] += 1
break
else:
counts["<no template>"] += 1
return counts