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