Clean tokenizer swap migration
This commit is contained in:
@@ -0,0 +1,166 @@
|
||||
#!/usr/bin/env python3
|
||||
import argparse
|
||||
import gzip
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
|
||||
def import_encoding(encoding_dir: Path):
|
||||
sys.path.insert(0, str(encoding_dir))
|
||||
import encoding_dsv4 # type: ignore
|
||||
|
||||
return encoding_dsv4
|
||||
|
||||
|
||||
def read_json(path):
|
||||
return json.loads(Path(path).read_text(encoding="utf-8"))
|
||||
|
||||
|
||||
def open_writer(path: Path):
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if path.suffix == ".gz":
|
||||
return gzip.open(path, "wt", encoding="utf-8")
|
||||
return path.open("w", encoding="utf-8")
|
||||
|
||||
|
||||
def quantiles(xs):
|
||||
if not xs:
|
||||
return {}
|
||||
xs = sorted(xs)
|
||||
return {
|
||||
"p50": xs[int((len(xs) - 1) * 0.50)],
|
||||
"p90": xs[int((len(xs) - 1) * 0.90)],
|
||||
"p95": xs[int((len(xs) - 1) * 0.95)],
|
||||
"p99": xs[int((len(xs) - 1) * 0.99)],
|
||||
"max": xs[-1],
|
||||
}
|
||||
|
||||
|
||||
def build_split(name, src_path, out_path, tok, enc, cutoff_len):
|
||||
rows = read_json(src_path)
|
||||
stats = {
|
||||
"split": name,
|
||||
"source": str(src_path),
|
||||
"output": str(out_path),
|
||||
"n": len(rows),
|
||||
"cutoff_len": cutoff_len,
|
||||
"truncated": 0,
|
||||
"prompt_tokens": [],
|
||||
"response_tokens": [],
|
||||
"total_tokens": [],
|
||||
"eos_in_labels": 0,
|
||||
"prefix_mismatch": 0,
|
||||
}
|
||||
|
||||
with open_writer(out_path) as f:
|
||||
for idx, row in enumerate(rows):
|
||||
instruction = (row.get("instruction") or "").strip()
|
||||
output = row.get("output") or ""
|
||||
messages_prompt = [{"role": "user", "content": instruction}]
|
||||
messages_full = [
|
||||
{"role": "user", "content": instruction},
|
||||
{"role": "assistant", "content": output},
|
||||
]
|
||||
prompt_text = enc.encode_messages(messages_prompt, thinking_mode="chat")
|
||||
full_text = enc.encode_messages(messages_full, thinking_mode="chat")
|
||||
if not full_text.startswith(prompt_text):
|
||||
stats["prefix_mismatch"] += 1
|
||||
|
||||
prompt_ids = tok(prompt_text, add_special_tokens=False).input_ids
|
||||
full_ids = tok(full_text, add_special_tokens=False).input_ids
|
||||
response_ids = full_ids[len(prompt_ids) :]
|
||||
labels = [-100] * len(prompt_ids) + response_ids
|
||||
|
||||
truncated = False
|
||||
if len(full_ids) > cutoff_len:
|
||||
truncated = True
|
||||
full_ids = full_ids[:cutoff_len]
|
||||
labels = labels[:cutoff_len]
|
||||
stats["truncated"] += 1
|
||||
|
||||
eos_in_labels = tok.eos_token_id in [x for x in labels if x != -100]
|
||||
stats["eos_in_labels"] += int(eos_in_labels)
|
||||
stats["prompt_tokens"].append(len(prompt_ids))
|
||||
stats["response_tokens"].append(len(response_ids))
|
||||
stats["total_tokens"].append(len(prompt_ids) + len(response_ids))
|
||||
|
||||
out = {
|
||||
"id": row.get("id", f"{name}_{idx:07d}"),
|
||||
"split": name,
|
||||
"source_task": row.get("task_type") or row.get("task") or row.get("category"),
|
||||
"source": row.get("source"),
|
||||
"thinking_mode": "chat",
|
||||
"messages": messages_full,
|
||||
"prompt_text": prompt_text,
|
||||
"full_text": full_text,
|
||||
"prompt_tokens": len(prompt_ids),
|
||||
"response_tokens": len(response_ids),
|
||||
"total_tokens": len(prompt_ids) + len(response_ids),
|
||||
"truncated": truncated,
|
||||
"eos_in_labels": eos_in_labels,
|
||||
"input_ids": full_ids,
|
||||
"labels": labels,
|
||||
}
|
||||
f.write(json.dumps(out, ensure_ascii=False) + "\n")
|
||||
|
||||
for key in ["prompt_tokens", "response_tokens", "total_tokens"]:
|
||||
stats[key] = quantiles(stats[key])
|
||||
stats["truncated_rate"] = stats["truncated"] / max(1, stats["n"])
|
||||
stats["eos_label_rate"] = stats["eos_in_labels"] / max(1, stats["n"])
|
||||
return stats
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--base-dir", default="/ssd/yi/Tokenizer_Swap")
|
||||
p.add_argument("--tokenizer", default="model_building/generated_models/Qwen3-0.6B-DSV4-tokenizer-remap-v2")
|
||||
p.add_argument("--encoding-dir", default="external/deepseek_v4_encoding")
|
||||
p.add_argument("--data-dir", required=True)
|
||||
p.add_argument("--out-dir", required=True)
|
||||
p.add_argument("--train-file", required=True)
|
||||
p.add_argument("--validation-file", default="fixed_validation.json")
|
||||
p.add_argument("--case-file", default="fixed_case.json")
|
||||
p.add_argument("--cutoff-len", type=int, default=2048)
|
||||
p.add_argument("--gzip", action="store_true")
|
||||
args = p.parse_args()
|
||||
|
||||
base = Path(args.base_dir)
|
||||
data_dir = base / args.data_dir
|
||||
out_dir = base / args.out_dir
|
||||
enc = import_encoding(base / args.encoding_dir)
|
||||
tok = AutoTokenizer.from_pretrained(base / args.tokenizer, trust_remote_code=True)
|
||||
suffix = ".jsonl.gz" if args.gzip else ".jsonl"
|
||||
split_files = {
|
||||
"train": args.train_file,
|
||||
"validation": args.validation_file,
|
||||
"case": args.case_file,
|
||||
}
|
||||
all_stats = {
|
||||
"tokenizer": str(base / args.tokenizer),
|
||||
"encoding_dir": str(base / args.encoding_dir),
|
||||
"data_dir": str(data_dir),
|
||||
"cutoff_len": args.cutoff_len,
|
||||
"eos_token": tok.eos_token,
|
||||
"eos_token_id": tok.eos_token_id,
|
||||
"splits": {},
|
||||
}
|
||||
for split, filename in split_files.items():
|
||||
stats = build_split(
|
||||
split,
|
||||
data_dir / filename,
|
||||
out_dir / f"{split}_dsv4_chat_tokenized{suffix}",
|
||||
tok,
|
||||
enc,
|
||||
args.cutoff_len,
|
||||
)
|
||||
all_stats["splits"][split] = stats
|
||||
print(json.dumps(stats, ensure_ascii=False), flush=True)
|
||||
(out_dir / "build_stats.json").write_text(json.dumps(all_stats, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
print(out_dir / "build_stats.json")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user