svd-code / sdg /preprocessing /preprocess_ot4_959k_code.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
3.82 kB
"""
Preprocess hero_run_4_code (code domain).
Source: https://huggingface.co/datasets/mlfoundations-dev/hero_run_4_code
Dataset stats:
- 959,024 rows
- Deduplicated by `instruction_seed`
This script deduplicates by instruction_seed (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 = "mlfoundations-dev/hero_run_4_code"
EXPECTED_UNIQUE_PROMPTS = 9_168
HF_REPO_ID = "teetone/ot4-code-deduped"
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 instruction_seed.
# 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
for row in ds:
total += 1
prompt = row.get("instruction_seed")
if prompt is None:
continue
rec = {
"conversations": [
{"from": "human", "value": prompt},
{"from": "gpt", "value": row["output"]},
],
}
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 total % 100_000 == 0:
print(f" Processed {total:,} rows, {len(prompt_reservoir):,} unique prompts ...")
print(f" Total rows streamed: {total:,}")
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:
prompt = rec["conversations"][0]["value"]
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()