svd-code / sdg /preprocessing /preprocess_nemotron_cascade_science.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
34.4 kB
"""
Preprocess Nemotron-Cascade-2-SFT-Data (science subset).
Source: https://huggingface.co/datasets/nvidia/Nemotron-Cascade-2-SFT-Data
Each row has top-level columns `domain`, `source`, `messages` (list of
{role, content} dicts: system, user, assistant), and `generator`. This script:
- Extracts the user-role message as the prompt.
- Deduplicates by user prompt via reservoir sampling (one random rollout
per unique prompt).
- Keeps the top-level `source` column (e.g. "Nemotron-Cascade-1",
"Nemotron-Science-v1").
- Emits records in the `openthoughts4` conversations format so SDG configs
can consume the dataset directly.
"""
from __future__ import annotations
import argparse
import json
import os
import random
import sys
import time
from pathlib import Path
# ── HF offline mode ──────────────────────────────────────────────────────
# Must run BEFORE `from datasets import ...` because datasets/huggingface_hub
# read these env vars into module-level constants at import time. Forcing them
# here makes a cached run do ZERO network calls (no revision check that could
# hang/fail when offline).
# --offline : force offline.
# --no-offline : force online (overrides everything).
# Default mode is download-to-cache (NOT streaming), so offline auto-enables
# IFF the dataset is already cached — the first download still needs network,
# and explicit --stream always stays online. A full (non --dry-run) run uploads
# to HF and therefore MUST stay online, so offline is only auto-enabled for dry
# runs (explicit --offline is still honored everywhere).
def _resolve_offline_flags() -> None:
argv = sys.argv[1:]
if "--no-offline" in argv:
return
want_offline = "--offline" in argv
streaming = "--stream" in argv
is_dry_run = "--dry-run" in argv
if not want_offline and is_dry_run and not streaming:
hf_home = os.environ.get("HF_HOME") or os.path.expanduser("~/.cache/huggingface")
# Cache folder name HF derives from "nvidia/Nemotron-Cascade-2-SFT-Data".
cache_dir = Path(hf_home) / "hub" / "datasets--nvidia--Nemotron-Cascade-2-SFT-Data"
want_offline = cache_dir.exists()
if want_offline:
os.environ.setdefault("HF_DATASETS_OFFLINE", "1")
os.environ.setdefault("HF_HUB_OFFLINE", "1")
_resolve_offline_flags()
from datasets import load_dataset
from huggingface_hub import HfApi, create_repo
from tqdm import tqdm
_REPO_ROOT = Path(__file__).resolve().parents[2]
if str(_REPO_ROOT) not in sys.path:
sys.path.insert(0, str(_REPO_ROOT))
from sdg.preprocessing.dedupe import (
MinHashDeduplicator,
SemanticDeduplicator,
find_common_prefixes,
strip_template,
template_hit_stats,
)
DATASET_ID = "nvidia/Nemotron-Cascade-2-SFT-Data"
DATASET_SUBSET = "science"
HF_REPO_ID = "teetone/nemotron-cascade-2-science-deduped"
EXPECTED_UNIQUE_PROMPTS: int | None = None
def extract_role(messages: list[dict], role: str) -> str | None:
"""Return the content of the first message matching `role`, or None."""
for m in messages:
if m.get("role") == role:
content = m.get("content")
if isinstance(content, str):
return content
return None
def _write_dedup_audit_section(
f,
stage_name: str,
threshold: float,
records_before: list[dict],
keep_indices: list[int],
clusters: dict,
max_text_chars: int = 600,
) -> None:
"""Append one stage's per-cluster audit info to file handle f.
`clusters` is the {rep_idx -> [member_idx, ...]} dict returned by
dedup(..., return_clusters=True). We show each non-singleton cluster with
its KEPT representative followed by the DROP'd members so a human can
eyeball *why* the algorithm merged them.
"""
n_in = len(records_before)
n_out = len(keep_indices)
non_singleton = {rep: m for rep, m in clusters.items() if len(m) > 1}
f.write(f"\n{'=' * 78}\n")
f.write(f"STAGE: {stage_name} (threshold={threshold})\n")
f.write(f" Input: {n_in:,} records\n")
f.write(f" Output: {n_out:,} records\n")
f.write(f" Removed: {n_in - n_out:,} records\n")
f.write(f" Clusters of size >= 2: {len(non_singleton):,}\n")
f.write(f"{'=' * 78}\n\n")
if not non_singleton:
f.write("(No clusters of size >= 2 — nothing was deduped at this stage.)\n")
return
# Sort by cluster size descending so the most aggressive merges show first.
sorted_clusters = sorted(non_singleton.items(), key=lambda kv: -len(kv[1]))
for cluster_idx, (rep, members) in enumerate(sorted_clusters, 1):
f.write(f"[CLUSTER {cluster_idx}] size={len(members)}\n")
for idx in members:
rec = records_before[idx]
prompt = rec["conversations"][0]["value"]
response = rec["conversations"][1]["value"]
marker = "KEPT" if idx == rep else "DROP"
if len(prompt) > max_text_chars:
shown = prompt[:max_text_chars] + f"...[+{len(prompt) - max_text_chars} chars]"
else:
shown = prompt
f.write(
f" {marker} [idx={idx:>6} resp_len={len(response):>6,}]\n"
f" {shown!r}\n"
)
f.write("\n")
def preprocess_and_upload(
seed: int = 42,
dry_run: bool = False,
stream: bool = False,
limit: int | None = None,
audit_path: str | None = None,
dump_embeddings: str | None = None,
auto_strip_templates: bool = True,
template_min_count: int = 10,
skip_minhash: bool = False,
minhash_threshold: float = 0.8,
minhash_num_perm: int = 128,
minhash_shingle_size: int = 5,
skip_semantic: bool = False,
semantic_threshold: float = 0.95,
embed_model: str = "BAAI/bge-small-en-v1.5",
embed_batch_size: int = 128,
semantic_topk: int = 10,
device: str = "auto",
) -> None:
random.seed(seed)
mode = "streaming" if stream else "download-to-cache"
offline = os.environ.get("HF_HUB_OFFLINE") == "1"
print(f"Loading {DATASET_ID} (subset={DATASET_SUBSET}, {mode}, "
f"offline={'on' if offline else 'off'}) ...")
ds = load_dataset(DATASET_ID, DATASET_SUBSET, split="train", streaming=stream)
prompt_reservoir: dict[str, tuple[dict, int]] = {}
total = 0
skipped_no_user = 0
skipped_no_assistant = 0
pbar = tqdm(ds, desc="Streaming Nemotron-Cascade-2 (science)", unit=" rows", smoothing=0.05)
for row in pbar:
total += 1
if limit is not None and total > limit:
total -= 1 # we counted this one but won't process it
break
messages = row.get("messages") or []
user_prompt = extract_role(messages, "user")
if not user_prompt:
skipped_no_user += 1
pbar.set_postfix(unique=len(prompt_reservoir), no_user=skipped_no_user, no_asst=skipped_no_assistant)
continue
assistant_response = extract_role(messages, "assistant")
if assistant_response is None:
skipped_no_assistant += 1
pbar.set_postfix(unique=len(prompt_reservoir), no_user=skipped_no_user, no_asst=skipped_no_assistant)
continue
source = row.get("source")
rec = {
"conversations": [
{"from": "human", "value": user_prompt},
{"from": "gpt", "value": assistant_response},
],
"source": source,
}
if user_prompt not in prompt_reservoir:
prompt_reservoir[user_prompt] = (rec, 1)
else:
_, count = prompt_reservoir[user_prompt]
count += 1
if random.randint(1, count) == 1:
prompt_reservoir[user_prompt] = (rec, count)
else:
prompt_reservoir[user_prompt] = (prompt_reservoir[user_prompt][0], count)
pbar.set_postfix(unique=len(prompt_reservoir), no_user=skipped_no_user, no_asst=skipped_no_assistant)
pbar.close()
print(f" Total rows streamed: {total:,}")
print(f" Skipped (no user message): {skipped_no_user:,}")
print(f" Skipped (no assistant message): {skipped_no_assistant:,}")
print(f" Unique user prompts: {len(prompt_reservoir):,}")
records = [row for row, _ in prompt_reservoir.values()]
n_after_exact = len(records)
print(f" Exact-deduped rows: {n_after_exact:,}")
# Resolve audit file path: --limit auto-enables auditing unless --skip-audit-on-limit.
write_audit = audit_path is not None or limit is not None
audit_handle = None
if write_audit:
if audit_path is None:
ts = time.strftime("%Y%m%d_%H%M%S")
audit_path = f"output/logs/dedupe_audit_{ts}.txt"
Path(audit_path).parent.mkdir(parents=True, exist_ok=True)
audit_handle = open(audit_path, "w")
audit_handle.write(
f"Dedup audit for nemotron-cascade-2-science (limit={limit}, dry_run={dry_run})\n"
f"Pipeline: exact reservoir -> MinHash -> Semantic\n"
f"After exact dedup: {n_after_exact:,} unique prompts\n"
)
audit_handle.write(f"\n{'=' * 78}\n")
audit_handle.write("RUN CONFIG\n")
audit_handle.write(f"{'=' * 78}\n")
audit_handle.write(f" seed = {seed}\n")
audit_handle.write(f" stream = {stream}\n")
audit_handle.write(f" offline = {offline}\n")
audit_handle.write(f" limit = {limit}\n")
audit_handle.write(f" dry_run = {dry_run}\n")
audit_handle.write(f" auto_strip_templates = {auto_strip_templates}\n")
audit_handle.write(f" template_min_count = {template_min_count}\n")
audit_handle.write(f" skip_minhash = {skip_minhash}\n")
audit_handle.write(f" minhash_threshold = {minhash_threshold}\n")
audit_handle.write(f" minhash_num_perm = {minhash_num_perm}\n")
audit_handle.write(f" minhash_shingle_size = {minhash_shingle_size}\n")
audit_handle.write(f" skip_semantic = {skip_semantic}\n")
audit_handle.write(f" semantic_threshold = {semantic_threshold}\n")
audit_handle.write(f" embed_model = {embed_model}\n")
audit_handle.write(f" embed_batch_size = {embed_batch_size}\n")
audit_handle.write(f" semantic_topk = {semantic_topk}\n")
audit_handle.write(f" device = {device}\n")
audit_handle.flush()
print(f"\n[audit] Writing per-cluster audit to: {audit_path}")
# ── Optional: auto-detect templates and strip before similarity comp ──
# Original prompts in `records` are NEVER mutated; stripping only affects
# the prompts passed to MinHash / Semantic for similarity computation.
templates: list[str] = []
if auto_strip_templates:
print(
f"\n[template-strip] Auto-detecting common prefix templates "
f"(min_count={template_min_count}) ..."
)
raw_prompts = [r["conversations"][0]["value"] for r in records]
templates = find_common_prefixes(raw_prompts, min_count=template_min_count)
if templates:
hit_counts = template_hit_stats(raw_prompts, templates)
print(f" Detected {len(templates)} template(s):")
for tmpl in templates:
preview = tmpl[:80].replace("\n", "\\n")
ellipsis = "..." if len(tmpl) > 80 else ""
print(
f" [{hit_counts[tmpl]:>6,} hits, {len(tmpl):>4} chars] "
f"{preview!r}{ellipsis}"
)
print(f" [{hit_counts['<no template>']:>6,} prompts unchanged]")
if audit_handle is not None:
audit_handle.write(f"\n{'=' * 78}\n")
audit_handle.write(
f"TEMPLATE STRIPPING (auto-detected, min_count={template_min_count})\n"
)
audit_handle.write(f" Found {len(templates)} template(s):\n")
audit_handle.write(f"{'=' * 78}\n\n")
for tmpl in templates:
audit_handle.write(
f"[{hit_counts[tmpl]:>6,} hits, {len(tmpl):>4} chars]\n"
f" {tmpl!r}\n\n"
)
audit_handle.write(
f"[{hit_counts['<no template>']:>6,} prompts unchanged by stripping]\n"
)
audit_handle.flush()
else:
print(f" No common templates found at min_count={template_min_count}.")
# ── Stage 2: MinHash near-duplicate dedup ─────────────────────────
if skip_minhash:
print("\nSkipping MinHash near-dup stage (--skip-minhash)")
n_after_minhash = n_after_exact
else:
print(
f"\n── MinHash near-dup "
f"(threshold={minhash_threshold}, num_perm={minhash_num_perm}, "
f"shingle_size={minhash_shingle_size}) ──"
)
records_before_minhash = records
prompts = [
strip_template(r["conversations"][0]["value"], templates) for r in records
]
rep_key = lambda i: -len(records[i]["conversations"][1]["value"])
result = MinHashDeduplicator(
threshold=minhash_threshold,
num_perm=minhash_num_perm,
shingle_size=minhash_shingle_size,
seed=seed,
).dedup(prompts, key_fn=rep_key, return_clusters=write_audit)
if write_audit:
keep, clusters_mh = result
else:
keep, clusters_mh = result, None
records = [records[i] for i in keep]
n_after_minhash = len(records)
print(
f" MinHash kept {n_after_minhash:,} / {n_after_exact:,} "
f"({n_after_exact - n_after_minhash:,} removed)"
)
if write_audit and audit_handle is not None:
_write_dedup_audit_section(
audit_handle,
stage_name=f"MinHash near-dup (num_perm={minhash_num_perm}, "
f"shingle_size={minhash_shingle_size})",
threshold=minhash_threshold,
records_before=records_before_minhash,
keep_indices=keep,
clusters=clusters_mh,
)
audit_handle.flush()
# ── Stage 3: Semantic paraphrase dedup ────────────────────────────
if skip_semantic:
print("\nSkipping semantic dedup stage (--skip-semantic)")
n_after_semantic = n_after_minhash
else:
print(
f"\n── Semantic dedup "
f"(model={embed_model}, threshold={semantic_threshold}) ──"
)
records_before_semantic = records
prompts = [
strip_template(r["conversations"][0]["value"], templates) for r in records
]
rep_key = lambda i: -len(records[i]["conversations"][1]["value"])
deduper = SemanticDeduplicator(
model_name=embed_model,
threshold=semantic_threshold,
batch_size=embed_batch_size,
device=device,
topk=semantic_topk,
)
if dump_embeddings is not None:
# Persist the (expensive) embeddings + per-record metadata so the
# post-MinHash set can be re-clustered at any threshold and analyzed
# (topic coverage, threshold sweep) offline without re-embedding.
import numpy as np
Path(dump_embeddings).mkdir(parents=True, exist_ok=True)
embeddings = deduper.encode(prompts)
np.save(str(Path(dump_embeddings) / "embeddings.npy"), embeddings)
with open(Path(dump_embeddings) / "records.jsonl", "w") as _df:
for r in records:
_df.write(json.dumps({
"prompt": r["conversations"][0]["value"],
"stripped": strip_template(r["conversations"][0]["value"], templates),
"resp_len": len(r["conversations"][1]["value"]),
"source": r.get("source"),
}, ensure_ascii=False) + "\n")
print(f" [dump-embeddings] saved {len(records):,} embeddings + records to {dump_embeddings}")
result = deduper.dedup_from_embeddings(
embeddings, key_fn=rep_key, return_clusters=write_audit
)
else:
result = deduper.dedup(prompts, key_fn=rep_key, return_clusters=write_audit)
if write_audit:
keep, clusters_sem = result
else:
keep, clusters_sem = result, None
records = [records[i] for i in keep]
n_after_semantic = len(records)
print(
f" Semantic kept {n_after_semantic:,} / {n_after_minhash:,} "
f"({n_after_minhash - n_after_semantic:,} removed)"
)
if write_audit and audit_handle is not None:
_write_dedup_audit_section(
audit_handle,
stage_name=f"Semantic paraphrase (model={embed_model}, topk={semantic_topk})",
threshold=semantic_threshold,
records_before=records_before_semantic,
keep_indices=keep,
clusters=clusters_sem,
)
audit_handle.flush()
if audit_handle is not None:
audit_handle.close()
print(f"\n[audit] Wrote per-cluster dedup audit to: {audit_path}")
# ── Assertions ────────────────────────────────────────────────────
if EXPECTED_UNIQUE_PROMPTS is not None:
assert len(records) == EXPECTED_UNIQUE_PROMPTS, (
f"Expected {EXPECTED_UNIQUE_PROMPTS} rows, got {len(records)}"
)
seen_prompts: set[str] = set()
for rec in records:
prompt = rec["conversations"][0]["value"]
assert prompt not in seen_prompts, f"Duplicate prompt found: {prompt[:80]}..."
seen_prompts.add(prompt)
print(" All assertions passed.")
# ── Source distribution ──────────────────────────────────────────
source_counts: dict[str, int] = {}
for rec in records:
src = rec.get("source") or "<unknown>"
source_counts[src] = source_counts.get(src, 0) + 1
print()
print("──────── Source breakdown (final dedup'd records) ────────")
for src, n in sorted(source_counts.items(), key=lambda x: -x[1]):
print(f" {src}: {n:,}")
print("──────────────────────────────────────────────────────────")
# ── Multiple-choice detection ────────────────────────────────────
mc_count = 0
non_mc_count = 0
for rec in records:
prompt = rec["conversations"][0]["value"]
if "multiple choice" in prompt.lower():
mc_count += 1
else:
non_mc_count += 1
pct_mc = (mc_count / len(records) * 100) if records else 0.0
print()
print("──────── Multiple-choice breakdown ────────")
print(f" Contains 'multiple choice': {mc_count:,} ({pct_mc:.1f}%)")
print(f" Does NOT contain 'multiple choice': {non_mc_count:,} ({100 - pct_mc:.1f}%)")
print("───────────────────────────────────────────")
# ── Sample preview ────────────────────────────────────────────────
if dry_run:
n_preview = min(5, len(records))
print()
print(f"──────── Sample preview ({n_preview} records that WOULD upload) ────────")
for i, rec in enumerate(random.sample(records, n_preview)):
print(f"\n[Sample {i + 1}/{n_preview}]")
blob = json.dumps(rec, ensure_ascii=False, indent=2)
print(blob[:2000])
if len(blob) > 2000:
print(" ... (truncated to 2000 chars)")
print("─────────────────────────────────────────────────────────────────")
# ── Write final jsonl + upload to HuggingFace ─────────────────────
# Write to a stable local path FIRST so a failed/blocked upload never throws
# away the (expensive) pipeline output — it can be re-uploaded from this file
# without re-running the whole pipeline.
out_path = Path("output/final") / (HF_REPO_ID.split("/")[-1] + ".jsonl")
out_path.parent.mkdir(parents=True, exist_ok=True)
print(f"\n Writing {len(records):,} final records to {out_path} ...")
t0 = time.perf_counter()
with open(out_path, "w") as f:
for rec in tqdm(records, desc=" jsonl write", unit=" rec", smoothing=0.05):
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
file_size_mb = out_path.stat().st_size / 1e6
print(f" wrote {file_size_mb:,.1f} MB in {time.perf_counter() - t0:.1f}s")
if dry_run:
print("[dry-run] Skipping upload to HuggingFace.")
else:
print(f"\nUploading {file_size_mb:,.1f} MB to {HF_REPO_ID} ...")
api = HfApi()
create_repo(HF_REPO_ID, repo_type="dataset", exist_ok=True)
t0 = time.perf_counter()
api.upload_file(
path_or_fileobj=str(out_path),
path_in_repo="data/train-00000-of-00001.jsonl",
repo_id=HF_REPO_ID,
repo_type="dataset",
)
print(f" uploaded in {time.perf_counter() - t0:.1f}s")
print()
print("──────── Pipeline attrition ────────")
print(f" {'Stage':<38s} {'Kept':>12s} {'Removed':>12s} {'% kept vs prev':>17s}")
def _attrition_row(label: str, kept: int, prev: int | None) -> None:
if prev is None:
print(f" {label:<38s} {kept:>12,} {'-':>12s} {'-':>17s}")
return
removed = prev - kept
pct = (kept / prev * 100) if prev else 0.0
print(f" {label:<38s} {kept:>12,} {removed:>12,} {pct:>16.2f}%")
_attrition_row("Streamed from HF", total, None)
_attrition_row("After exact dedup (reservoir)", n_after_exact, total)
_attrition_row("After MinHash near-dup", n_after_minhash, n_after_exact)
_attrition_row("After semantic dedup", n_after_semantic, n_after_minhash)
print("──────────────────────────────────────────────────────────────────────────────────")
print()
print("──────── Final stats ────────")
print(f" Total rows streamed: {total:,}")
print(f" Skipped (no user message): {skipped_no_user:,}")
print(f" Skipped (no assistant message): {skipped_no_assistant:,}")
print(f" Unique user prompts (final): {len(records):,}")
print(f" Contains 'multiple choice': {mc_count:,} ({pct_mc:.1f}%)")
print(f" Does NOT contain 'multiple choice': {non_mc_count:,} ({100 - pct_mc:.1f}%)")
if dry_run:
print(f" Rows that WOULD upload: {len(records):,} (dry-run, not uploaded)")
print(f" HF dataset (target): https://huggingface.co/datasets/{HF_REPO_ID}")
else:
print(f" Rows uploaded to HF: {len(records):,}")
print(f" HF dataset: https://huggingface.co/datasets/{HF_REPO_ID}")
print("─────────────────────────────")
def sample_and_upload(
n_sample: int,
seed: int = 42,
dry_run: bool = False,
source_path: str | None = None,
target_repo: str | None = None,
) -> None:
"""Shuffle + randomly sample N records from the already-deduped dataset and
upload to a separate HF repo. Independent of the dedup pipeline.
Source is the local final jsonl produced by a full run
(output/final/<repo>.jsonl) if present, else the HF dataset HF_REPO_ID.
Records are copied verbatim, so the output format is identical to the source.
"""
random.seed(seed)
if target_repo is None:
suffix = f"{n_sample // 1000}k" if n_sample % 1000 == 0 else str(n_sample)
target_repo = f"{HF_REPO_ID}-{suffix}"
local_final = (
Path(source_path) if source_path
else Path("output/final") / (HF_REPO_ID.split("/")[-1] + ".jsonl")
)
# Reservoir sampling: a uniform random sample of N in a single pass, holding
# only N lines in memory (no need to count the source first).
if local_final.exists():
print(f"Sampling {n_sample:,} from local file: {local_final}")
reservoir: list[str] = []
total = 0
with open(local_final) as f:
for i, line in enumerate(f):
total += 1
if i < n_sample:
reservoir.append(line)
else:
j = random.randint(0, i)
if j < n_sample:
reservoir[j] = line
else:
print(f"Local file not found; loading from HF: {HF_REPO_ID}")
ds = load_dataset(HF_REPO_ID, split="train")
total = len(ds)
if n_sample > total:
raise ValueError(f"--sample {n_sample} exceeds source size {total:,}")
idx = random.sample(range(total), n_sample)
reservoir = [json.dumps(ds[i], ensure_ascii=False) + "\n" for i in idx]
print(f" source records: {total:,}")
if n_sample > total:
raise ValueError(f"--sample {n_sample} exceeds source size {total:,}")
random.shuffle(reservoir)
out_path = Path("output/final") / (target_repo.split("/")[-1] + ".jsonl")
out_path.parent.mkdir(parents=True, exist_ok=True)
with open(out_path, "w") as f:
f.writelines(reservoir)
mb = out_path.stat().st_size / 1e6
print(f" wrote {len(reservoir):,} sampled records ({mb:,.1f} MB) to {out_path}")
src_counts: dict[str, int] = {}
for line in reservoir:
s = json.loads(line).get("source") or "<unknown>"
src_counts[s] = src_counts.get(s, 0) + 1
print(" sample source breakdown:",
{k: f"{v:,}" for k, v in sorted(src_counts.items(), key=lambda x: -x[1])})
if dry_run:
print(f"[dry-run] Skipping upload to {target_repo}.")
return
print(f"\nUploading {mb:,.1f} MB to {target_repo} ...")
create_repo(target_repo, repo_type="dataset", exist_ok=True)
t0 = time.perf_counter()
HfApi().upload_file(
path_or_fileobj=str(out_path),
path_in_repo="data/train-00000-of-00001.jsonl",
repo_id=target_repo,
repo_type="dataset",
)
print(f" uploaded in {time.perf_counter() - t0:.1f}s "
f"-> https://huggingface.co/datasets/{target_repo}")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--dry-run",
action="store_true",
help="Run preprocessing and assertions but skip the HuggingFace upload.",
)
parser.add_argument("--seed", type=int, default=42, help="Random seed for reservoir sampling.")
parser.add_argument("--stream", action=argparse.BooleanOptionalAction, default=False,
help="Stream the dataset over the network instead of downloading it to "
"the local HF cache first. Default: OFF (download-to-cache, then "
"iterate from disk — faster and re-runnable offline). Use --stream "
"to opt into streaming, or --no-stream to force download mode.")
parser.add_argument("--offline", action="store_true",
help="Force HF offline mode (HF_DATASETS_OFFLINE=1, HF_HUB_OFFLINE=1): "
"read dataset + embed model purely from cache, zero network calls. "
"Handled before importing datasets; listed here for --help/validation.")
parser.add_argument("--no-offline", action="store_true",
help="Force online mode (e.g. to refresh the cache). In download mode, "
"offline auto-enables once cached; this overrides that behavior.")
# ── Sampling mode (separate code path; does NOT run the dedup pipeline) ──
parser.add_argument("--sample", type=int, default=None, metavar="N",
help="SAMPLING MODE: shuffle + randomly sample N records from the already-"
f"deduped dataset ({HF_REPO_ID}) and upload to a separate repo. "
"Reads output/final/<repo>.jsonl if present, else loads from HF. "
"Does not run the dedup pipeline. Honors --seed and --dry-run.")
parser.add_argument("--sample-repo", type=str, default=None,
help="Target HF repo for --sample (default: '<HF_REPO_ID>-<N>k').")
parser.add_argument("--sample-source", type=str, default=None,
help="Override the source jsonl for --sample (default: "
"output/final/<repo>.jsonl).")
parser.add_argument("--limit", type=int, default=None,
help="Process at most N rows from the stream (debug aid). "
"When set, also writes a per-cluster audit file by default.")
parser.add_argument("--audit-path", type=str, default=None,
help="Path for the per-cluster dedup audit file. Auto-generated "
"under output/logs/ when --limit is set; specify here to override "
"or to enable audit on a full run.")
parser.add_argument("--dump-embeddings", type=str, default=None,
help="Directory to save the semantic-stage embeddings.npy + records.jsonl "
"(post-MinHash set). Lets you re-cluster at any threshold and run "
"topic/coverage analysis offline without re-embedding.")
parser.add_argument("--auto-strip-templates", action=argparse.BooleanOptionalAction,
default=True,
help="Auto-detect and strip frequent prompt-template prefixes before "
"computing similarity. Original prompts in records are unchanged; "
"only the dedup signal sees stripped versions. Default: ON (prevents "
"boilerplate-driven over-merging). Use --no-auto-strip-templates "
"to disable.")
parser.add_argument("--template-min-count", type=int, default=10,
help="Min number of prompts a prefix must appear in to be treated "
"as a template (default 10).")
parser.add_argument("--skip-minhash", action="store_true", help="Skip MinHash near-dup stage.")
parser.add_argument("--minhash-threshold", type=float, default=0.8,
help="Jaccard threshold for MinHash (default 0.8).")
parser.add_argument("--minhash-num-perm", type=int, default=128,
help="Number of MinHash permutations (default 128).")
parser.add_argument("--minhash-shingle-size", type=int, default=5,
help="Word-shingle size for MinHash (default 5).")
parser.add_argument("--skip-semantic", action="store_true", help="Skip semantic dedup stage.")
parser.add_argument("--semantic-threshold", type=float, default=0.95,
help="Cosine threshold for semantic dedup (default 0.95). 0.92 was found "
"to over-merge distinct STEM questions into topic blobs; 0.95 keeps "
"~109K more records and removes ~70%% fewer from large clusters.")
parser.add_argument("--embed-model", type=str, default="BAAI/bge-small-en-v1.5",
help="Sentence-transformers model for embedding.")
parser.add_argument("--embed-batch-size", type=int, default=128, help="Embedding batch size.")
parser.add_argument("--semantic-topk", type=int, default=10,
help="FAISS top-k neighbors per query.")
parser.add_argument("--device", type=str, default="auto",
choices=["auto", "mps", "cuda", "cpu"], help="Embedding device.")
args = parser.parse_args()
if args.sample is not None:
sample_and_upload(
n_sample=args.sample,
seed=args.seed,
dry_run=args.dry_run,
source_path=args.sample_source,
target_repo=args.sample_repo,
)
sys.exit(0)
preprocess_and_upload(
seed=args.seed,
dry_run=args.dry_run,
stream=args.stream,
limit=args.limit,
audit_path=args.audit_path,
dump_embeddings=args.dump_embeddings,
auto_strip_templates=args.auto_strip_templates,
template_min_count=args.template_min_count,
skip_minhash=args.skip_minhash,
minhash_threshold=args.minhash_threshold,
minhash_num_perm=args.minhash_num_perm,
minhash_shingle_size=args.minhash_shingle_size,
skip_semantic=args.skip_semantic,
semantic_threshold=args.semantic_threshold,
embed_model=args.embed_model,
embed_batch_size=args.embed_batch_size,
semantic_topk=args.semantic_topk,
device=args.device,
)