Clean tokenizer swap migration
This commit is contained in:
Executable
+208
@@ -0,0 +1,208 @@
|
||||
#!/usr/bin/env python3
|
||||
import argparse
|
||||
import gzip
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
SCIENCE_KEYWORDS = re.compile(
|
||||
r"\b(physics|chemistry|biology|biological|medical|medicine|clinical|anatomy|physiology|genetics|ecology|astronomy|geology|neuroscience|experiment|hypothesis|laboratory|scientific|research|disease|protein|cell|molecule|atom|energy|force|gravity|electric|magnetic|quantum|planet|star|organism|evolution|climate|weather|ecosystem|bacteria|virus|vaccine|enzyme|hormone|blood|brain|heart|lung|kidney|cancer|therapy|diagnosis)\b",
|
||||
re.I,
|
||||
)
|
||||
|
||||
|
||||
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 first_user_assistant(messages):
|
||||
user = None
|
||||
assistant = None
|
||||
for msg in messages or []:
|
||||
role = msg.get("role")
|
||||
content = clean_text(msg.get("content"))
|
||||
if not content:
|
||||
continue
|
||||
if role == "user" and user is None:
|
||||
user = content
|
||||
elif role == "assistant" and user is not None:
|
||||
assistant = content
|
||||
break
|
||||
return user, assistant
|
||||
|
||||
|
||||
def qa_text(user, assistant):
|
||||
if not user or not assistant:
|
||||
return ""
|
||||
return f"Question:\n{user}\n\nAnswer:\n{assistant}"
|
||||
|
||||
|
||||
def extract_row_text(row):
|
||||
if isinstance(row, dict) and row.get("messages"):
|
||||
u, a = first_user_assistant(row.get("messages"))
|
||||
if u and a:
|
||||
return qa_text(u, a)
|
||||
for key in ("text", "content", "markdown", "raw_content"):
|
||||
if isinstance(row, dict):
|
||||
text = clean_text(row.get(key))
|
||||
if text:
|
||||
return text
|
||||
return ""
|
||||
|
||||
|
||||
def write_record(writer, stats, source, text, ntok, metadata=None):
|
||||
idx = stats["docs"]
|
||||
writer.write(json.dumps({
|
||||
"id": f"science_fix_{idx:09d}",
|
||||
"category": "science",
|
||||
"source": source,
|
||||
"text": text,
|
||||
"token_count": ntok,
|
||||
"metadata": metadata or {},
|
||||
}, ensure_ascii=False) + "\n")
|
||||
stats["docs"] += 1
|
||||
stats["tokens"] += ntok
|
||||
stats["tokens_by_source"][source] += ntok
|
||||
|
||||
|
||||
def copy_existing(existing, writer, target):
|
||||
tokens = 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 add_jsonl(path, source, tok, writer, stats, args, capability_filter=None):
|
||||
with open(path, "r", encoding="utf-8", errors="replace") as f:
|
||||
for line in f:
|
||||
if stats["tokens"] >= args.target_tokens:
|
||||
break
|
||||
if not line.strip():
|
||||
continue
|
||||
try:
|
||||
row = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if capability_filter and row.get("capability") not in capability_filter:
|
||||
continue
|
||||
text = extract_row_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
|
||||
write_record(writer, stats, source, text, ntok, {k: row.get(k) for k in ("id", "capability", "source_id", "split") if k in row})
|
||||
if stats["docs"] % args.log_every == 0:
|
||||
print(json.dumps({"event": "progress", "docs": stats["docs"], "tokens": stats["tokens"], "target": args.target_tokens, "source": source, "elapsed_sec": time.time() - stats["started_at"]}, ensure_ascii=False), flush=True)
|
||||
|
||||
|
||||
def add_fineweb_edu_parquets(paths, tok, writer, stats, args):
|
||||
import pyarrow.parquet as pq
|
||||
for path in paths:
|
||||
if stats["tokens"] >= args.target_tokens:
|
||||
break
|
||||
source = f"fineweb_edu_science_supplement:{path.name}"
|
||||
before = stats["tokens"]
|
||||
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_row_text(row)
|
||||
if not text or not SCIENCE_KEYWORDS.search(text[:12000]):
|
||||
stats["rejected"]["not_science_like"] += 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
|
||||
write_record(writer, stats, source, text, ntok, {k: row.get(k) for k in ("url", "dump", "language", "language_score") if k in row})
|
||||
if stats["docs"] % args.log_every == 0:
|
||||
print(json.dumps({"event": "progress", "docs": stats["docs"], "tokens": stats["tokens"], "target": args.target_tokens, "source": source, "elapsed_sec": time.time() - stats["started_at"]}, ensure_ascii=False), flush=True)
|
||||
print(json.dumps({"event": "source_done", "source": source, "tokens_added": stats["tokens"] - before, "total_tokens": stats["tokens"]}, ensure_ascii=False), flush=True)
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--tokenizer", default="model_building/generated_models/Qwen3-0.6B-DSV4-tokenizer-remap-v2")
|
||||
ap.add_argument("--out-dir", default="data/cpt_docmix_5b_sources_8192_20260614")
|
||||
ap.add_argument("--target-tokens", type=int, default=150_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("--parquet-batch-size", type=int, default=1000)
|
||||
args = ap.parse_args()
|
||||
|
||||
base = Path.cwd()
|
||||
out_dir = base / args.out_dir
|
||||
doc_dir = out_dir / "documents"
|
||||
final_out = doc_dir / "science.jsonl.gz"
|
||||
tmp_out = doc_dir / "science.jsonl.gz.tmp"
|
||||
if tmp_out.exists():
|
||||
tmp_out.unlink()
|
||||
|
||||
tok = AutoTokenizer.from_pretrained(base / args.tokenizer, trust_remote_code=True)
|
||||
stats = {"target_tokens": args.target_tokens, "tokens": 0, "docs": 0, "tokens_by_source": Counter(), "rejected": Counter(), "started_at": time.time(), "sources": []}
|
||||
|
||||
with gzip.open(tmp_out, "wt", encoding="utf-8") as writer:
|
||||
t, d = copy_existing(final_out, writer, args.target_tokens)
|
||||
stats["tokens"] += t; stats["docs"] += d; stats["tokens_by_source"]["existing_science_docmix"] += t
|
||||
print(json.dumps({"event": "copied_existing", "docs": d, "tokens": t}, ensure_ascii=False), flush=True)
|
||||
|
||||
mix = base / "data/training_mix_v4_train1m_test2p8k_noupsample_nobbh_20260611/train_1m.jsonl"
|
||||
if mix.exists() and stats["tokens"] < args.target_tokens:
|
||||
before = stats["tokens"]
|
||||
add_jsonl(mix, "training_mix_v4_science_logic", tok, writer, stats, args, {"science_reasoning", "logic"})
|
||||
stats["sources"].append({"source": str(mix), "tokens": stats["tokens"] - before})
|
||||
|
||||
paths = sorted((base / "data/raw_parquets/fineweb_edu").glob("*.parquet"))
|
||||
if paths and stats["tokens"] < args.target_tokens:
|
||||
add_fineweb_edu_parquets(paths, tok, writer, stats, args)
|
||||
|
||||
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 / "science_5b_fix_stats.json").write_text(json.dumps(stats, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
(out_dir / ".science_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