File size: 4,516 Bytes
58258b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
"""
Preprocess OpenThoughts3-1.2M (math domain only).

Source: https://huggingface.co/datasets/open-thoughts/OpenThoughts3-1.2M

Dataset stats (math domain):
  - 850,000 rows from source 'ai2-adapt-dev/openmath-2-math'
  - 53,125 unique user prompts (~16 rollouts per prompt)
  - Conversations: 2 turns (human prompt + gpt response with <think> reasoning)

This script filters the math domain, deduplicates by user prompt
(picking one random rollout per prompt via reservoir sampling),
asserts correctness, and uploads to HuggingFace.
"""

from __future__ import annotations

import json
import random
import tempfile
from pathlib import Path

from datasets import load_dataset
from huggingface_hub import HfApi, create_repo


DATASET_ID = "open-thoughts/OpenThoughts3-1.2M"
EXPECTED_UNIQUE_PROMPTS = 53_125
HF_REPO_ID = "teetone/openthoughts3-math-deduped-53K"


def extract_prompt(conversations: list[dict]) -> str | None:
    """Return the first human turn from a conversation."""
    for turn in conversations:
        if turn["from"] == "human":
            return turn["value"]
    return None


def preprocess_and_upload(seed: int = 42) -> None:
    random.seed(seed)

    print(f"Loading {DATASET_ID} (streaming) ...")
    ds = load_dataset(DATASET_ID, split="train", streaming=True)

    # Single streaming pass: reservoir-sample one row per unique prompt.
    # For each prompt we track (chosen_row, count_seen) so we can do
    # uniform random selection without storing all rollouts.
    prompt_reservoir: dict[str, tuple[dict, int]] = {}
    total = 0
    math_count = 0

    for row in ds:
        total += 1
        if row["domain"] != "math":
            continue
        math_count += 1

        prompt = extract_prompt(row["conversations"])
        if prompt is None:
            continue

        rec = {
            "difficulty": row["difficulty"],
            "source": row["source"],
            "domain": row["domain"],
            "conversations": row["conversations"],
        }

        if prompt not in prompt_reservoir:
            prompt_reservoir[prompt] = (rec, 1)
        else:
            _, count = prompt_reservoir[prompt]
            count += 1
            # Reservoir sampling: replace with probability 1/count
            if random.randint(1, count) == 1:
                prompt_reservoir[prompt] = (rec, count)
            else:
                prompt_reservoir[prompt] = (prompt_reservoir[prompt][0], count)

        if math_count % 100_000 == 0:
            print(f"  Processed {math_count:,} math rows, {len(prompt_reservoir):,} unique prompts ...")

    print(f"  Total rows streamed: {total:,}")
    print(f"  Math rows: {math_count:,}")
    print(f"  Unique prompts: {len(prompt_reservoir):,}")

    records = [row for row, _ in prompt_reservoir.values()]
    print(f"  Deduplicated rows: {len(records):,}")

    # ── Assertions ────────────────────────────────────────────────────
    assert len(records) == EXPECTED_UNIQUE_PROMPTS, (
        f"Expected {EXPECTED_UNIQUE_PROMPTS} rows, got {len(records)}"
    )

    seen_prompts: set[str] = set()
    for rec in records:
        assert rec["domain"] == "math", f"Non-math row found: domain={rec['domain']}"
        prompt = extract_prompt(rec["conversations"])
        assert prompt not in seen_prompts, f"Duplicate prompt found: {prompt[:80]}..."
        seen_prompts.add(prompt)

    assert len(seen_prompts) == EXPECTED_UNIQUE_PROMPTS, (
        f"Expected {EXPECTED_UNIQUE_PROMPTS} unique prompts, got {len(seen_prompts)}"
    )
    print("  All assertions passed.")

    # ── Upload to HuggingFace ─────────────────────────────────────────
    print(f"Uploading to {HF_REPO_ID} ...")
    api = HfApi()
    create_repo(HF_REPO_ID, repo_type="dataset", exist_ok=True)

    with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f:
        temp_path = f.name
        for rec in records:
            f.write(json.dumps(rec, ensure_ascii=False) + "\n")

    api.upload_file(
        path_or_fileobj=temp_path,
        path_in_repo="data/train-00000-of-00001.jsonl",
        repo_id=HF_REPO_ID,
        repo_type="dataset",
    )
    Path(temp_path).unlink()

    print(f"Done! Dataset available at: https://huggingface.co/datasets/{HF_REPO_ID}")


if __name__ == "__main__":
    preprocess_and_upload()