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