MultimodalReasoning3B / training /tensor_cache.py
gaaaaaaaaaaa's picture
Initial commit before training
61449ba verified
Raw History Blame Contribute Delete
5.24 kB
from __future__ import annotations
import hashlib
import importlib.util
import json
from pathlib import Path
from typing import Any, Iterable, Sequence
CACHE_VERSION = 1
def _json_hash(value: Any) -> str:
encoded = json.dumps(value, ensure_ascii=False, sort_keys=True).encode("utf-8")
return hashlib.sha256(encoded).hexdigest()
def _hash_text(value: Any) -> str:
return hashlib.sha256(str(value or "").encode("utf-8")).hexdigest()
def parquet_reader_available() -> bool:
return any(
importlib.util.find_spec(package) is not None
for package in ("pyarrow", "polars", "pandas")
)
def dataset_source_files(dataset_name: str, data_dir: str | Path) -> list[Path]:
root = Path(data_dir)
if dataset_name == "multimodal":
return sorted((root / "MultimodalReasoning").glob("*.jsonl"))
if dataset_name == "reasoning_dpo":
return sorted((root / "ReasoningDPO").glob("*.jsonl"))
if dataset_name == "limo":
return [root / "limo.jsonl"]
if dataset_name == "complex_bespoke":
return [root / "ComplexReasoningBespoke.parquet"]
if dataset_name == "all":
files: list[Path] = []
files.extend(dataset_source_files("multimodal", root))
files.extend(dataset_source_files("reasoning_dpo", root))
files.extend(dataset_source_files("limo", root))
files.extend(dataset_source_files("complex_bespoke", root))
return files
return []
def file_metadata(path: Path) -> dict[str, Any]:
if not path.exists():
return {
"path": str(path),
"exists": False,
}
stat = path.stat()
return {
"path": str(path),
"exists": True,
"size": stat.st_size,
"mtime_ns": stat.st_mtime_ns,
}
def data_digest(records: Sequence[dict[str, str]]) -> dict[str, Any]:
digest = hashlib.sha256()
for record in records:
digest.update(str(record.get("user", "")).encode("utf-8"))
digest.update(b"\0")
digest.update(str(record.get("assistant", "")).encode("utf-8"))
digest.update(b"\0\0")
return {
"rows": len(records),
"sha256": digest.hexdigest(),
}
def tokenizer_metadata(tokenizer: Any) -> dict[str, Any]:
return {
"name_or_path": str(getattr(tokenizer, "name_or_path", "")),
"class": tokenizer.__class__.__name__,
"eos_token": str(getattr(tokenizer, "eos_token", "")),
"pad_token": str(getattr(tokenizer, "pad_token", "")),
"padding_side": str(getattr(tokenizer, "padding_side", "")),
"chat_template_sha256": _hash_text(getattr(tokenizer, "chat_template", "")),
"vocab_size": len(tokenizer) if hasattr(tokenizer, "__len__") else None,
}
def build_cache_metadata(
*,
config: dict[str, Any],
tokenizer: Any,
provided_records: Sequence[dict[str, str]] | None = None,
) -> dict[str, Any]:
relevant_config = {
key: config.get(key)
for key in (
"model_name",
"data_dir",
"dataset_name",
"parquet_limit",
"system_prompt",
"dataset_format",
"append_eos_to_completion",
"max_samples",
"eval_split_size",
"max_eval_samples",
"shuffle_data",
"seed",
"max_length",
"completion_only_loss",
"assistant_only_loss",
)
}
metadata: dict[str, Any] = {
"cache_version": CACHE_VERSION,
"config": relevant_config,
"tokenizer": tokenizer_metadata(tokenizer),
"parquet_reader_available": parquet_reader_available(),
}
if provided_records is not None:
metadata["provided_data"] = data_digest(provided_records)
else:
metadata["source_files"] = [
file_metadata(path)
for path in dataset_source_files(
str(config.get("dataset_name", "all")),
str(config.get("data_dir", "data")),
)
]
metadata["cache_key"] = _json_hash(metadata)[:24]
return metadata
def cache_path(cache_root: str | Path, metadata: dict[str, Any]) -> Path:
dataset_name = str(metadata["config"].get("dataset_name") or "dataset")
safe_name = "".join(
char if char.isalnum() or char in ("-", "_") else "_"
for char in dataset_name
)
return Path(cache_root) / f"{safe_name}_{metadata['cache_key']}"
def metadata_path(path: str | Path) -> Path:
return Path(path) / "metadata.json"
def read_metadata(path: str | Path) -> dict[str, Any] | None:
meta_path = metadata_path(path)
if not meta_path.exists():
return None
with meta_path.open("r", encoding="utf-8") as f:
return json.load(f)
def write_metadata(path: str | Path, metadata: dict[str, Any]) -> None:
meta_path = metadata_path(path)
meta_path.parent.mkdir(parents=True, exist_ok=True)
with meta_path.open("w", encoding="utf-8") as f:
json.dump(metadata, f, ensure_ascii=False, indent=2, sort_keys=True)
def metadata_matches(path: str | Path, metadata: dict[str, Any]) -> bool:
cached = read_metadata(path)
return cached == metadata