Files
model-training/llm/audit_cpt_corpus.py
T
2026-07-19 23:52:25 +04:00

153 lines
5.3 KiB
Python

"""Read-only audit of the materialized RU CPT corpus and training checkpoints.
This never rewrites the corpus. It verifies Arrow shape, block lengths/token
ranges, clean-source provenance, hashes, and the latest Trainer checkpoint.
Output is JSON so the result can be archived and compared across runs.
"""
from __future__ import annotations
import argparse
import hashlib
import json
import random
from pathlib import Path
from datasets import load_from_disk
from transformers import AutoTokenizer
from corpus_common import CLEAN, PACKED, has_mixed_script, cyrillic_ratio
MODEL = "Qwen/Qwen3-1.7B-Base"
BLOCK = 2048
def sha256(path: Path) -> str:
h = hashlib.sha256()
with path.open("rb") as f:
for chunk in iter(lambda: f.read(1024 * 1024), b""):
h.update(chunk)
return h.hexdigest()
def clean_stats() -> dict:
result = {}
for path in sorted(CLEAN.glob("*.jsonl")):
docs = chars = mixed = 0
cyr_sum = 0.0
with path.open(encoding="utf-8") as f:
for line in f:
if not line.strip():
continue
text = json.loads(line)["text"]
docs += 1
chars += len(text)
mixed += int(has_mixed_script(text))
cyr_sum += cyrillic_ratio(text)
result[path.stem] = {
"docs": docs,
"chars": chars,
"mean_doc_chars": round(chars / docs, 2) if docs else 0,
"mean_cyrillic_ratio": round(cyr_sum / docs, 6) if docs else 0,
"mixed_script_docs": mixed,
"bytes": path.stat().st_size,
"sha256": sha256(path),
}
return result
def checkpoint_stats(root: Path) -> dict:
states = []
for state_path in root.glob("checkpoint-*/trainer_state.json"):
state = json.loads(state_path.read_text(encoding="utf-8"))
states.append((int(state["global_step"]), state_path, state))
if not states:
return {"present": False}
step, state_path, state = max(states)
log = state.get("log_history", [])
last = next((x for x in reversed(log) if "loss" in x), {})
return {
"present": True,
"latest_checkpoint": str(state_path.parent),
"global_step": step,
"max_steps": state.get("max_steps"),
"progress_pct": round(100 * step / state["max_steps"], 3),
"epoch": state.get("epoch"),
"last_loss": last.get("loss"),
"last_grad_norm": last.get("grad_norm"),
"last_learning_rate": last.get("learning_rate"),
"full_weight_model_bytes": (state_path.parent / "model.safetensors").stat().st_size,
}
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--samples", type=int, default=64)
ap.add_argument("--seed", type=int, default=20260718)
ap.add_argument("--output")
ap.add_argument("--checkpoint-root",
default=str(Path(__file__).resolve().parent / "Qwen3-1.7B-ru-cpt"))
args = ap.parse_args()
ds = load_from_disk(str(PACKED))
if not len(ds):
raise SystemExit("packed dataset is empty")
rng = random.Random(args.seed)
indices = sorted(rng.sample(range(len(ds)), min(args.samples, len(ds))))
tokenizer = AutoTokenizer.from_pretrained(MODEL, local_files_only=True)
vocab = len(tokenizer)
bad_lengths = bad_ids = decoded_mixed = 0
samples = []
for idx in indices:
ids = ds[idx]["input_ids"]
bad_lengths += int(len(ids) != BLOCK)
bad_ids += sum(1 for token in ids if token < 0 or token >= vocab)
text = tokenizer.decode(ids, skip_special_tokens=True)
decoded_mixed += int(has_mixed_script(text))
samples.append({
"index": idx,
"tokens": len(ids),
"cyrillic_ratio": round(cyrillic_ratio(text), 6),
"mixed_script": has_mixed_script(text),
"preview": text[:240].replace("\n", " "),
})
arrows = sorted(PACKED.glob("*.arrow"))
clean = clean_stats()
total_chars = sum(x["chars"] for x in clean.values())
report = {
"schema_version": 1,
"read_only": True,
"model": MODEL,
"packed": {
"path": str(PACKED),
"rows": len(ds),
"block_tokens": BLOCK,
"total_tokens": len(ds) * BLOCK,
"arrow_bytes": sum(p.stat().st_size for p in arrows),
"arrow_sha256": {p.name: sha256(p) for p in arrows},
"sampled_rows": len(indices),
"sample_bad_lengths": bad_lengths,
"sample_out_of_vocab_ids": bad_ids,
"sample_mixed_script_rows": decoded_mixed,
},
"clean_sources": clean,
"clean_char_share_pct": {
key: round(100 * value["chars"] / total_chars, 3)
for key, value in clean.items()
},
"provenance_warning": (
"Packed Arrow has no per-block source/document metadata; source shares can "
"only be reconstructed from the clean inputs, not proven from Arrow alone."
),
"checkpoint": checkpoint_stats(Path(args.checkpoint_root)),
"samples": samples,
}
rendered = json.dumps(report, ensure_ascii=False, indent=2) + "\n"
if args.output:
Path(args.output).write_text(rendered, encoding="utf-8")
print(rendered, end="")
if __name__ == "__main__":
main()