Clean tokenizer swap migration
This commit is contained in:
Executable
+264
@@ -0,0 +1,264 @@
|
||||
#!/usr/bin/env python3
|
||||
import argparse
|
||||
import gzip
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
|
||||
from huggingface_hub import HfApi, hf_hub_download
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
MATHPILE_PRIORITY = [
|
||||
"train/textbooks/textbooks_markdown.jsonl.gz",
|
||||
"train/textbooks/synthetic_textbooks_markdown.jsonl.gz",
|
||||
"train/wikipedia/wikipedia_en_mathematics_nopic_2023-08_v0.2.jsonl.gz",
|
||||
"train/proofwiki/ProofWiki_definitions.jsonl.gz",
|
||||
"train/proofwiki/ProofWiki_theorem_proofs.jsonl.gz",
|
||||
"train/stackexchange/math.stackexchange.com.jsonl.gz",
|
||||
"train/stackexchange/mathoverflow.net.jsonl.gz",
|
||||
"train/stackexchange/physics.stackexchange.com.jsonl.gz",
|
||||
"train/arXiv/math_arXiv_v0.2_chunk_1.jsonl.gz",
|
||||
"train/arXiv/math_arXiv_v0.2_chunk_2.jsonl.gz",
|
||||
"train/arXiv/math_arXiv_v0.2_chunk_3.jsonl.gz",
|
||||
"train/arXiv/math_arXiv_v0.2_chunk_4.jsonl.gz",
|
||||
"train/commoncrawl/C4_math_docs_chunk_0.jsonl.gz",
|
||||
"train/commoncrawl/CC_math_docs_chunk_0.jsonl.gz",
|
||||
]
|
||||
|
||||
|
||||
def clean_text(text):
|
||||
if text is None:
|
||||
return ""
|
||||
text = str(text).replace("\x00", " ")
|
||||
text = re.sub(r"[ \t\r\f\v]+", " ", text)
|
||||
text = re.sub(r"\n{4,}", "\n\n\n", text)
|
||||
return text.strip()
|
||||
|
||||
|
||||
def extract_text(row):
|
||||
for key in ("text", "content", "markdown", "raw_content", "document"):
|
||||
if isinstance(row, dict):
|
||||
text = clean_text(row.get(key))
|
||||
if text:
|
||||
return text
|
||||
return ""
|
||||
|
||||
|
||||
def safe_source(text):
|
||||
return re.sub(r"[^A-Za-z0-9._-]+", "_", text)[:120]
|
||||
|
||||
|
||||
def download_with_retry(repo, filename, args):
|
||||
last = None
|
||||
for attempt in range(1, args.retries + 1):
|
||||
try:
|
||||
return hf_hub_download(
|
||||
repo_id=repo,
|
||||
repo_type="dataset",
|
||||
filename=filename,
|
||||
endpoint=args.endpoint,
|
||||
token=os.environ.get("HF_TOKEN"),
|
||||
local_dir=args.raw_dir,
|
||||
)
|
||||
except Exception as exc:
|
||||
last = exc
|
||||
print(json.dumps({"event": "download_retry", "repo": repo, "file": filename, "attempt": attempt, "error": repr(exc)[:800]}, ensure_ascii=False), flush=True)
|
||||
time.sleep(min(120, 5 * attempt))
|
||||
raise RuntimeError(f"download failed for {repo}:{filename}: {last!r}")
|
||||
|
||||
|
||||
def copy_existing(existing, writer, target):
|
||||
tokens = 0
|
||||
docs = 0
|
||||
if not existing.exists() or existing.stat().st_size < 1024:
|
||||
return tokens, docs
|
||||
with gzip.open(existing, "rt", encoding="utf-8", errors="replace") as f:
|
||||
for line in f:
|
||||
if not line.strip() or tokens >= target:
|
||||
continue
|
||||
try:
|
||||
rec = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
ntok = int(rec.get("token_count") or 0)
|
||||
if ntok <= 0:
|
||||
continue
|
||||
writer.write(json.dumps(rec, ensure_ascii=False) + "\n")
|
||||
tokens += ntok
|
||||
docs += 1
|
||||
return tokens, docs
|
||||
|
||||
|
||||
def iter_jsonl_gz(path):
|
||||
opener = gzip.open if str(path).endswith(".gz") else open
|
||||
with opener(path, "rt", encoding="utf-8", errors="replace") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
try:
|
||||
yield json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
|
||||
def add_jsonl_source(path, source_label, tok, writer, stats, args):
|
||||
for row in iter_jsonl_gz(path):
|
||||
if stats["tokens"] >= args.target_tokens:
|
||||
break
|
||||
text = extract_text(row)
|
||||
if not text:
|
||||
stats["rejected"]["empty"] += 1
|
||||
continue
|
||||
ntok = len(tok.encode(text, add_special_tokens=False))
|
||||
if ntok < args.min_tokens:
|
||||
stats["rejected"]["too_short"] += 1
|
||||
continue
|
||||
if ntok > args.max_doc_tokens:
|
||||
stats["rejected"]["too_long"] += 1
|
||||
continue
|
||||
idx = stats["docs"]
|
||||
writer.write(json.dumps({
|
||||
"id": f"math_fix_{idx:09d}",
|
||||
"category": "math",
|
||||
"source": source_label,
|
||||
"text": text,
|
||||
"token_count": ntok,
|
||||
"metadata": {k: row.get(k) for k in ("id", "url", "source", "title") if isinstance(row, dict) and k in row},
|
||||
}, ensure_ascii=False) + "\n")
|
||||
stats["docs"] += 1
|
||||
stats["tokens"] += ntok
|
||||
stats["tokens_by_source"][source_label] += ntok
|
||||
if stats["docs"] % args.log_every == 0:
|
||||
print(json.dumps({"event": "progress", "docs": stats["docs"], "tokens": stats["tokens"], "target": args.target_tokens, "source": source_label, "elapsed_sec": time.time() - stats["started_at"]}, ensure_ascii=False), flush=True)
|
||||
|
||||
|
||||
def add_parquet_source(path, source_label, tok, writer, stats, args):
|
||||
import pyarrow.parquet as pq
|
||||
pf = pq.ParquetFile(path)
|
||||
for batch in pf.iter_batches(batch_size=args.parquet_batch_size):
|
||||
if stats["tokens"] >= args.target_tokens:
|
||||
break
|
||||
for row in batch.to_pylist():
|
||||
if stats["tokens"] >= args.target_tokens:
|
||||
break
|
||||
text = extract_text(row)
|
||||
if not text:
|
||||
stats["rejected"]["empty"] += 1
|
||||
continue
|
||||
ntok = len(tok.encode(text, add_special_tokens=False))
|
||||
if ntok < args.min_tokens:
|
||||
stats["rejected"]["too_short"] += 1
|
||||
continue
|
||||
if ntok > args.max_doc_tokens:
|
||||
stats["rejected"]["too_long"] += 1
|
||||
continue
|
||||
idx = stats["docs"]
|
||||
writer.write(json.dumps({
|
||||
"id": f"math_fix_{idx:09d}",
|
||||
"category": "math",
|
||||
"source": source_label,
|
||||
"text": text,
|
||||
"token_count": ntok,
|
||||
"metadata": {k: row.get(k) for k in ("url", "date", "language", "language_score") if k in row},
|
||||
}, ensure_ascii=False) + "\n")
|
||||
stats["docs"] += 1
|
||||
stats["tokens"] += ntok
|
||||
stats["tokens_by_source"][source_label] += ntok
|
||||
if stats["docs"] % args.log_every == 0:
|
||||
print(json.dumps({"event": "progress", "docs": stats["docs"], "tokens": stats["tokens"], "target": args.target_tokens, "source": source_label, "elapsed_sec": time.time() - stats["started_at"]}, ensure_ascii=False), flush=True)
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--endpoint", default=os.environ.get("HF_ENDPOINT", "https://hf-mirror.com"))
|
||||
ap.add_argument("--tokenizer", default="model_building/generated_models/Qwen3-0.6B-DSV4-tokenizer-remap-v2")
|
||||
ap.add_argument("--raw-dir", default="data/raw_math_fix_20260614")
|
||||
ap.add_argument("--out-dir", default="data/cpt_docmix_5b_sources_8192_20260614")
|
||||
ap.add_argument("--target-tokens", type=int, default=500_000_000)
|
||||
ap.add_argument("--min-tokens", type=int, default=128)
|
||||
ap.add_argument("--max-doc-tokens", type=int, default=32768)
|
||||
ap.add_argument("--log-every", type=int, default=5000)
|
||||
ap.add_argument("--retries", type=int, default=16)
|
||||
ap.add_argument("--parquet-batch-size", type=int, default=1000)
|
||||
ap.add_argument("--keep-raw", action="store_true")
|
||||
ap.add_argument("--skip-mathpile", action="store_true")
|
||||
ap.add_argument("--skip-openwebmath-prefix-count", type=int, default=0)
|
||||
args = ap.parse_args()
|
||||
|
||||
base = Path.cwd()
|
||||
out_dir = base / args.out_dir
|
||||
doc_dir = out_dir / "documents"
|
||||
doc_dir.mkdir(parents=True, exist_ok=True)
|
||||
raw_dir = base / args.raw_dir
|
||||
raw_dir.mkdir(parents=True, exist_ok=True)
|
||||
args.raw_dir = str(raw_dir)
|
||||
|
||||
tok = AutoTokenizer.from_pretrained(base / args.tokenizer, trust_remote_code=True)
|
||||
final_out = doc_dir / "math.jsonl.gz"
|
||||
tmp_out = doc_dir / "math.jsonl.gz.tmp"
|
||||
if tmp_out.exists():
|
||||
tmp_out.unlink()
|
||||
|
||||
stats = {"target_tokens": args.target_tokens, "tokens": 0, "docs": 0, "tokens_by_source": Counter(), "rejected": Counter(), "sources": [], "started_at": time.time()}
|
||||
api = HfApi(endpoint=args.endpoint, token=os.environ.get("HF_TOKEN"))
|
||||
|
||||
with gzip.open(tmp_out, "wt", encoding="utf-8") as writer:
|
||||
existing_tokens, existing_docs = copy_existing(final_out, writer, args.target_tokens)
|
||||
stats["tokens"] += existing_tokens
|
||||
stats["docs"] += existing_docs
|
||||
stats["tokens_by_source"]["existing_math_docmix"] += existing_tokens
|
||||
print(json.dumps({"event": "copied_existing", "docs": existing_docs, "tokens": existing_tokens}, ensure_ascii=False), flush=True)
|
||||
|
||||
for filename in ([] if args.skip_mathpile else MATHPILE_PRIORITY):
|
||||
if stats["tokens"] >= args.target_tokens:
|
||||
break
|
||||
try:
|
||||
local = Path(download_with_retry("GAIR/MathPile", filename, args))
|
||||
before = stats["tokens"]
|
||||
add_jsonl_source(local, f"GAIR/MathPile:{filename}", tok, writer, stats, args)
|
||||
stats["sources"].append({"repo": "GAIR/MathPile", "file": filename, "tokens": stats["tokens"] - before})
|
||||
print(json.dumps({"event": "source_done", "source": f"GAIR/MathPile:{filename}", "tokens_added": stats["tokens"] - before, "total_tokens": stats["tokens"]}, ensure_ascii=False), flush=True)
|
||||
if not args.keep_raw:
|
||||
try: local.unlink()
|
||||
except Exception: pass
|
||||
except Exception as exc:
|
||||
stats["sources"].append({"repo": "GAIR/MathPile", "file": filename, "error": repr(exc)[:1000]})
|
||||
print(json.dumps({"event": "source_error", "source": f"GAIR/MathPile:{filename}", "error": repr(exc)[:1000]}, ensure_ascii=False), flush=True)
|
||||
|
||||
if stats["tokens"] < args.target_tokens:
|
||||
files = [f for f in api.list_repo_files("open-web-math/open-web-math", repo_type="dataset") if f.startswith("data/") and f.endswith(".parquet")]
|
||||
for filename in sorted(files)[args.skip_openwebmath_prefix_count:]:
|
||||
if stats["tokens"] >= args.target_tokens:
|
||||
break
|
||||
try:
|
||||
local = Path(download_with_retry("open-web-math/open-web-math", filename, args))
|
||||
before = stats["tokens"]
|
||||
add_parquet_source(local, f"open-web-math/open-web-math:{filename}", tok, writer, stats, args)
|
||||
stats["sources"].append({"repo": "open-web-math/open-web-math", "file": filename, "tokens": stats["tokens"] - before})
|
||||
print(json.dumps({"event": "source_done", "source": f"open-web-math/open-web-math:{filename}", "tokens_added": stats["tokens"] - before, "total_tokens": stats["tokens"]}, ensure_ascii=False), flush=True)
|
||||
if not args.keep_raw:
|
||||
try: local.unlink()
|
||||
except Exception: pass
|
||||
except Exception as exc:
|
||||
stats["sources"].append({"repo": "open-web-math/open-web-math", "file": filename, "error": repr(exc)[:1000]})
|
||||
print(json.dumps({"event": "source_error", "source": f"open-web-math/open-web-math:{filename}", "error": repr(exc)[:1000]}, ensure_ascii=False), flush=True)
|
||||
|
||||
if stats["tokens"] < args.target_tokens:
|
||||
raise SystemExit(f"only collected {stats['tokens']} / {args.target_tokens} tokens")
|
||||
|
||||
if final_out.exists():
|
||||
final_out.replace(final_out.with_suffix(".jsonl.gz.underfilled_20260614"))
|
||||
tmp_out.replace(final_out)
|
||||
stats["elapsed_sec"] = time.time() - stats["started_at"]
|
||||
stats["tokens_by_source"] = dict(stats["tokens_by_source"])
|
||||
stats["rejected"] = dict(stats["rejected"])
|
||||
(out_dir / "math_5b_fix_stats.json").write_text(json.dumps(stats, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
(out_dir / ".math_5b_ready").write_text(json.dumps({"tokens": stats["tokens"], "docs": stats["docs"], "elapsed_sec": stats["elapsed_sec"]}, ensure_ascii=False), encoding="utf-8")
|
||||
print(json.dumps({"event": "done", "tokens": stats["tokens"], "docs": stats["docs"], "elapsed_sec": stats["elapsed_sec"], "output": str(final_out)}, ensure_ascii=False), flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user